diff --git a/README.md b/README.md index 0f82f9f..7671e35 100644 --- a/README.md +++ b/README.md @@ -18,6 +18,32 @@ go install github.com/alyraffauf/tg/cmd/tg@latest ## Usage +### Authentication + +Log in interactively with OAuth: + +```bash +tg auth login alice.example.com +``` + +For headless use, pass an atproto app password as the second argument: + +```bash +tg auth login alice.example.com xxxx-xxxx-xxxx-xxxx +``` + +To avoid exposing the app password in shell history, pass it on standard input: + +```bash +printf '%s\n' "$ATPROTO_APP_PASSWORD" | tg auth login alice.example.com --password-stdin +``` + +Authentication is persisted locally. The current account is recorded in +`~/.config/tg/auth.json` (or `$XDG_CONFIG_HOME/tg/auth.json`); OAuth session +credentials are stored under `~/.config/tg/oauth/`, and app-password sessions +are stored in `~/.config/tg/password-session.json`. These files are created +with user-only permissions. Use `tg auth logout` to remove the active login. + `tg` auto-detects the repository from the `origin` remote when run inside a cloned Tangled repo. For now, only ssh origins are supported. You can also pass a fully-qualified `handle/repo` argument. ### Repositories diff --git a/atproto/auth.go b/atproto/auth.go index aaed245..e9f3def 100644 --- a/atproto/auth.go +++ b/atproto/auth.go @@ -3,17 +3,20 @@ package atproto import ( "context" "errors" + "fmt" "net/url" + "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/identity" "github.com/bluesky-social/indigo/atproto/syntax" "github.com/zalando/go-keyring" ) var ErrNotAuthenticated = errors.New("not authenticated") -// DefaultScopes are requested for a CLI session. The rpc scopes are needed for the -// PDS to mint service-auth JWTs for knot procedures. The blob scope is +// DefaultScopes are requested for a CLI session. The rpc scopes are needed for +// the PDS to mint service-auth JWTs for knot procedures. The blob scope is // required for uploading PR patch blobs to the PDS. var DefaultScopes = []string{ "atproto", @@ -51,13 +54,41 @@ func NewAuthManager(callbackURL string) *AuthManager { } } +// LoginWithPassword authenticates with an atproto app password and stores the +// resulting session in the keyring. Any existing OAuth session is cleared so +// only one auth method is active at a time. +func (m *AuthManager) LoginWithPassword(ctx context.Context, identifier, password string) error { + parsedIdentifier, err := syntax.ParseAtIdentifier(identifier) + if err != nil { + return err + } + persist := func(_ context.Context, data atclient.PasswordSessionData) { + _ = m.store.SavePasswordSession(context.Background(), data) + } + client, err := atclient.LoginWithPassword(ctx, identity.DefaultDirectory(), parsedIdentifier, password, "", persist) + if err != nil { + return err + } + passwordAuth, ok := client.Auth.(*atclient.PasswordAuth) + if !ok { + return errors.New("password login returned an unexpected auth type") + } + _ = m.store.DeleteSession(ctx, "", "") + return m.store.SavePasswordSession(ctx, passwordAuth.Session) +} + func (m *AuthManager) StartLogin(ctx context.Context, identifier string) (string, error) { return m.app.StartAuthFlow(ctx, identifier) } func (m *AuthManager) FinishLogin(ctx context.Context, query url.Values) error { _, err := m.app.ProcessCallback(ctx, query) - return err + if err != nil { + return err + } + // Clear any existing password session so only one auth method is active. + _ = m.store.DeletePasswordSession(ctx) + return nil } // CancelLogin cleans up any pending auth request written by StartLogin when the @@ -68,11 +99,21 @@ func (m *AuthManager) CancelLogin() { } func (m *AuthManager) CurrentDID(ctx context.Context) (syntax.DID, error) { - session, err := m.CurrentSession(ctx) + session, err := m.app.ResumeSession(ctx, "", "") + if err == nil { + return session.Data.AccountDID, nil + } + if !errors.Is(err, keyring.ErrNotFound) { + return "", err + } + passwordSession, err := m.store.GetPasswordSession(ctx) if err != nil { + if errors.Is(err, keyring.ErrNotFound) { + return "", ErrNotAuthenticated + } return "", err } - return session.Data.AccountDID, nil + return passwordSession.AccountDID, nil } func (m *AuthManager) CurrentSession(ctx context.Context) (*oauth.ClientSession, error) { @@ -86,22 +127,63 @@ func (m *AuthManager) CurrentSession(ctx context.Context) (*oauth.ClientSession, return session, nil } +// APIClient returns an API client and the account DID for the active session, +// whether OAuth or app-password. Token refreshes are persisted back to the +// keyring. +func (m *AuthManager) APIClient(ctx context.Context) (*atclient.APIClient, syntax.DID, error) { + session, err := m.app.ResumeSession(ctx, "", "") + if err == nil { + return session.APIClient(), session.Data.AccountDID, nil + } + if !errors.Is(err, keyring.ErrNotFound) { + return nil, "", err + } + passwordSession, err := m.store.GetPasswordSession(ctx) + if err != nil { + if errors.Is(err, keyring.ErrNotFound) { + return nil, "", ErrNotAuthenticated + } + return nil, "", err + } + persist := func(_ context.Context, data atclient.PasswordSessionData) { + _ = m.store.SavePasswordSession(context.Background(), data) + } + client := atclient.ResumePasswordSession(*passwordSession, persist) + return client, passwordSession.AccountDID, nil +} + func (m *AuthManager) Logout(ctx context.Context) error { err := m.app.Logout(ctx, "", "") - if err == nil { + switch { + case err == nil: return nil + case errors.Is(err, keyring.ErrNotFound): + // No OAuth session; continue to password logout below. + default: + // Corrupt or transient OAuth failure — force clear so the user can + // re-login instead of being locked out. + if deleteErr := m.store.DeleteSession(ctx, "", ""); deleteErr == nil { + return nil + } + return err } - if errors.Is(err, keyring.ErrNotFound) { - return ErrNotAuthenticated + + passwordSession, err := m.store.GetPasswordSession(ctx) + if err != nil { + if errors.Is(err, keyring.ErrNotFound) { + return ErrNotAuthenticated + } + return err } - // Logout failed partway — most commonly because the session entry is - // corrupt or the keyring is transiently unavailable. The upstream Logout - // aborts before DeleteSession in that case, so the bad entry would - // otherwise be unrecoverable short of manual keyring surgery. Best-effort - // clear it so the user can re-login; if that also fails, surface the - // original error. - if deleteErr := m.store.DeleteSession(ctx, "", ""); deleteErr == nil { + client := atclient.ResumePasswordSession(*passwordSession, nil) + passwordAuth, ok := client.Auth.(*atclient.PasswordAuth) + if !ok { + // Corrupt password session — force clear. + _ = m.store.DeletePasswordSession(ctx) return nil } - return err + if err := passwordAuth.Logout(ctx, client.Client); err != nil { + return fmt.Errorf("revoke password session: %w", err) + } + return m.store.DeletePasswordSession(ctx) } diff --git a/atproto/auth_test.go b/atproto/auth_test.go new file mode 100644 index 0000000..5346a1b --- /dev/null +++ b/atproto/auth_test.go @@ -0,0 +1,153 @@ +package atproto + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/bluesky-social/indigo/atproto/atclient" + "github.com/bluesky-social/indigo/atproto/syntax" + "github.com/zalando/go-keyring" +) + +func samplePasswordSession(did syntax.DID, host string) atclient.PasswordSessionData { + return atclient.PasswordSessionData{ + AccountDID: did, + AccessToken: "access", + RefreshToken: "refresh", + Host: host, + } +} + +func TestPasswordSessionRoundTrip(t *testing.T) { + store := testKeyringStore(newFakeKeyring()) + ctx := context.Background() + did := mustDID(t, "did:plc:aaaabbbbccccddddeeeeffff") + + want := samplePasswordSession(did, "https://pds.example") + if err := store.SavePasswordSession(ctx, want); err != nil { + t.Fatalf("SavePasswordSession: %v", err) + } + + got, err := store.GetPasswordSession(ctx) + if err != nil { + t.Fatalf("GetPasswordSession: %v", err) + } + if got.AccountDID != want.AccountDID { + t.Errorf("AccountDID = %q, want %q", got.AccountDID, want.AccountDID) + } + if got.AccessToken != want.AccessToken { + t.Errorf("AccessToken = %q, want %q", got.AccessToken, want.AccessToken) + } + if got.RefreshToken != want.RefreshToken { + t.Errorf("RefreshToken = %q, want %q", got.RefreshToken, want.RefreshToken) + } + if got.Host != want.Host { + t.Errorf("Host = %q, want %q", got.Host, want.Host) + } +} + +func TestPasswordSessionNotFound(t *testing.T) { + store := testKeyringStore(newFakeKeyring()) + _, err := store.GetPasswordSession(context.Background()) + if !errors.Is(err, keyring.ErrNotFound) { + t.Errorf("GetPasswordSession = %v, want keyring.ErrNotFound", err) + } +} + +func TestCurrentDIDFallsBackToPassword(t *testing.T) { + store := testKeyringStore(newFakeKeyring()) + manager := newAuthManagerForTest("http://127.0.0.1:8095/callback", store) + ctx := context.Background() + did := mustDID(t, "did:plc:aaaabbbbccccddddeeeeffff") + + if err := store.SavePasswordSession(ctx, samplePasswordSession(did, "https://pds.example")); err != nil { + t.Fatalf("SavePasswordSession: %v", err) + } + + got, err := manager.CurrentDID(ctx) + if err != nil { + t.Fatalf("CurrentDID: %v", err) + } + if got != did { + t.Errorf("CurrentDID = %q, want %q", got, did) + } +} + +func TestCurrentDIDNotAuthenticated(t *testing.T) { + manager := newAuthManagerForTest("http://127.0.0.1:8095/callback", testKeyringStore(newFakeKeyring())) + _, err := manager.CurrentDID(context.Background()) + if !errors.Is(err, ErrNotAuthenticated) { + t.Errorf("CurrentDID = %v, want ErrNotAuthenticated", err) + } +} + +func TestAPIClientReturnsPasswordClient(t *testing.T) { + store := testKeyringStore(newFakeKeyring()) + manager := newAuthManagerForTest("http://127.0.0.1:8095/callback", store) + ctx := context.Background() + did := mustDID(t, "did:plc:aaaabbbbccccddddeeeeffff") + + if err := store.SavePasswordSession(ctx, samplePasswordSession(did, "https://pds.example")); err != nil { + t.Fatalf("SavePasswordSession: %v", err) + } + + client, gotDID, err := manager.APIClient(ctx) + if err != nil { + t.Fatalf("APIClient: %v", err) + } + if gotDID != did { + t.Errorf("DID = %q, want %q", gotDID, did) + } + if _, ok := client.Auth.(*atclient.PasswordAuth); !ok { + t.Errorf("Auth = %T, want *atclient.PasswordAuth", client.Auth) + } +} + +func TestAPIClientNotAuthenticated(t *testing.T) { + manager := newAuthManagerForTest("http://127.0.0.1:8095/callback", testKeyringStore(newFakeKeyring())) + _, _, err := manager.APIClient(context.Background()) + if !errors.Is(err, ErrNotAuthenticated) { + t.Errorf("APIClient = %v, want ErrNotAuthenticated", err) + } +} + +// TestPasswordLogoutRevokesAndRemovesSession verifies that logging out of a +// password session revokes it at the PDS (using the refresh token) and then +// deletes it from the keyring. +func TestPasswordLogoutRevokesAndRemovesSession(t *testing.T) { + var authorization string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/xrpc/com.atproto.server.deleteSession" { + t.Errorf("unexpected path %q", r.URL.Path) + } + if r.Method != http.MethodPost { + t.Errorf("unexpected method %q", r.Method) + } + authorization = r.Header.Get("Authorization") + w.WriteHeader(http.StatusOK) + })) + defer server.Close() + + store := testKeyringStore(newFakeKeyring()) + manager := newAuthManagerForTest("http://127.0.0.1:8095/callback", store) + ctx := context.Background() + did := mustDID(t, "did:plc:aaaabbbbccccddddeeeeffff") + + if err := store.SavePasswordSession(ctx, samplePasswordSession(did, server.URL)); err != nil { + t.Fatalf("SavePasswordSession: %v", err) + } + + if err := manager.Logout(ctx); err != nil { + t.Fatalf("Logout: %v", err) + } + + if authorization != "Bearer refresh" { + t.Errorf("Authorization = %q, want %q", authorization, "Bearer refresh") + } + if _, err := store.GetPasswordSession(ctx); !errors.Is(err, keyring.ErrNotFound) { + t.Errorf("password session should have been removed, got: %v", err) + } +} diff --git a/atproto/keyring_store.go b/atproto/keyring_store.go index ee64d9f..b5afb30 100644 --- a/atproto/keyring_store.go +++ b/atproto/keyring_store.go @@ -7,6 +7,7 @@ import ( "fmt" "sync" + "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/syntax" "github.com/zalando/go-keyring" @@ -56,6 +57,8 @@ func NewKeyringStore() *KeyringStore { const currentSessionKey = "session:current" +const currentPasswordKey = "password:current" + func requestKey(state string) string { return "request:" + state } @@ -110,6 +113,28 @@ func (s *KeyringStore) DeleteSession(_ context.Context, _ syntax.DID, _ string) return s.deleteSecret(currentSessionKey) } +func (s *KeyringStore) GetPasswordSession(_ context.Context) (*atclient.PasswordSessionData, error) { + s.mu.Lock() + defer s.mu.Unlock() + var session atclient.PasswordSessionData + if err := s.getSecret(currentPasswordKey, &session); err != nil { + return nil, err + } + return &session, nil +} + +func (s *KeyringStore) SavePasswordSession(_ context.Context, session atclient.PasswordSessionData) error { + s.mu.Lock() + defer s.mu.Unlock() + return s.saveSecret(currentPasswordKey, session) +} + +func (s *KeyringStore) DeletePasswordSession(_ context.Context) error { + s.mu.Lock() + defer s.mu.Unlock() + return s.deleteSecret(currentPasswordKey) +} + func (s *KeyringStore) GetAuthRequestInfo(_ context.Context, state string) (*oauth.AuthRequestData, error) { s.mu.Lock() defer s.mu.Unlock() diff --git a/internal/cli/api.go b/internal/cli/api.go index 967bdf0..b5623af 100644 --- a/internal/cli/api.go +++ b/internal/cli/api.go @@ -3,11 +3,13 @@ package cli import ( "bytes" "encoding/json" + "errors" "fmt" "io" "net/http" "strings" + "github.com/alyraffauf/tg/atproto" "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/syntax" "github.com/spf13/cobra" @@ -39,11 +41,14 @@ var apiCmd = &cobra.Command{ return fmt.Errorf("method must be GET or POST, got %q", apiMethod) } - session, err := requireAuthSession(cmd.Context()) + client, _, err := auth.APIClient(cmd.Context()) if err != nil { - return err + if errors.Is(err, atproto.ErrNotAuthenticated) { + return fmt.Errorf("not logged in; run \"tg auth login\" first") + } + return fmt.Errorf("resume auth session: %w", err) } - response, err := doAPIRequest(cmd, session.APIClient(), endpoint, method, fields) + response, err := doAPIRequest(cmd, client, endpoint, method, fields) if err != nil { return err } diff --git a/internal/cli/atproto_auth.go b/internal/cli/atproto_auth.go index 1fff1d3..5a148a8 100644 --- a/internal/cli/atproto_auth.go +++ b/internal/cli/atproto_auth.go @@ -6,24 +6,15 @@ import ( "fmt" "github.com/alyraffauf/tg/atproto" - "github.com/bluesky-social/indigo/atproto/auth/oauth" ) -func requireAuthSession(ctx context.Context) (*oauth.ClientSession, error) { - session, err := auth.CurrentSession(ctx) +func authenticatedATProto(ctx context.Context) (*atproto.ATProto, string, error) { + client, did, err := auth.APIClient(ctx) if err != nil { if errors.Is(err, atproto.ErrNotAuthenticated) { - return nil, fmt.Errorf("not logged in; run \"tg auth login\" first") + return nil, "", fmt.Errorf("not logged in; run \"tg auth login\" first") } - return nil, fmt.Errorf("resume OAuth session: %w", err) - } - return session, nil -} - -func authenticatedATProto(ctx context.Context) (*atproto.ATProto, string, error) { - session, err := requireAuthSession(ctx) - if err != nil { - return nil, "", err + return nil, "", fmt.Errorf("resume auth session: %w", err) } - return &atproto.ATProto{Client: session.APIClient()}, session.Data.AccountDID.String(), nil + return &atproto.ATProto{Client: client}, did.String(), nil } diff --git a/internal/cli/auth_login.go b/internal/cli/auth_login.go index 45bef48..c26330f 100644 --- a/internal/cli/auth_login.go +++ b/internal/cli/auth_login.go @@ -3,26 +3,40 @@ package cli import ( "context" "fmt" + "io" "net/http" "os" "os/exec" "runtime" + "strings" "github.com/spf13/cobra" ) +var authLoginPasswordStdin bool + var authLoginCmd = &cobra.Command{ - Use: "login [handle]", - Short: "Log in to atproto via OAuth", - Long: `Log in to atproto via OAuth using a local browser callback.`, - Args: cobra.MaximumNArgs(1), + Use: "login [app-password]", + Short: "Log in to atproto via OAuth or an app password", + Long: `Log in with OAuth, or use an app password as the second argument for headless login.`, + Args: cobra.RangeArgs(1, 2), RunE: func(cmd *cobra.Command, args []string) error { - identifier := "" - if len(args) == 1 { - identifier = args[0] + identifier := args[0] + password, usePassword, err := loginPassword(args, authLoginPasswordStdin, cmd.InOrStdin()) + if err != nil { + return err } - if identifier == "" { - return fmt.Errorf("handle or DID required") + if usePassword { + if err := auth.LoginWithPassword(cmd.Context(), identifier, password); err != nil { + return err + } + did, err := auth.CurrentDID(cmd.Context()) + if err != nil { + fmt.Fprintln(os.Stderr, "Login completed but session could not be confirmed.") + return err + } + fmt.Printf("Logged in as %s\n", did) + return nil } server, resultChannel, err := runCallbackServer() @@ -62,6 +76,31 @@ var authLoginCmd = &cobra.Command{ }, } +func init() { + authLoginCmd.Flags().BoolVar(&authLoginPasswordStdin, "password-stdin", false, "Read the app password from standard input") +} + +func loginPassword(args []string, fromStdin bool, stdin io.Reader) (string, bool, error) { + if !fromStdin { + if len(args) < 2 { + return "", false, nil + } + return args[1], true, nil + } + if len(args) == 2 { + return "", false, fmt.Errorf("app password argument and --password-stdin cannot be used together") + } + data, err := io.ReadAll(stdin) + if err != nil { + return "", false, fmt.Errorf("read app password from stdin: %w", err) + } + password := strings.TrimSpace(string(data)) + if password == "" { + return "", false, fmt.Errorf("app password from stdin is empty") + } + return password, true, nil +} + func runCallbackServer() (*http.Server, <-chan error, error) { resultChannel := make(chan error, 1) diff --git a/internal/cli/auth_login_test.go b/internal/cli/auth_login_test.go new file mode 100644 index 0000000..c11b183 --- /dev/null +++ b/internal/cli/auth_login_test.go @@ -0,0 +1,36 @@ +package cli + +import ( + "strings" + "testing" +) + +func TestLoginPassword(t *testing.T) { + tests := []struct { + name string + args []string + stdinFlag bool + stdin string + want string + wantUse bool + wantErr bool + }{ + {"oauth", []string{"alice.example"}, false, "", "", false, false}, + {"argument", []string{"alice.example", "app-pass"}, false, "", "app-pass", true, false}, + {"stdin", []string{"alice.example"}, true, "app-pass\n", "app-pass", true, false}, + {"both", []string{"alice.example", "app-pass"}, true, "other", "", false, true}, + {"empty stdin", []string{"alice.example"}, true, "\n", "", false, true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, use, err := loginPassword(tt.args, tt.stdinFlag, strings.NewReader(tt.stdin)) + if (err != nil) != tt.wantErr { + t.Fatalf("error = %v, wantErr %v", err, tt.wantErr) + } + if got != tt.want || use != tt.wantUse { + t.Fatalf("got (%q, %v), want (%q, %v)", got, use, tt.want, tt.wantUse) + } + }) + } +} diff --git a/internal/cli/auth_token.go b/internal/cli/auth_token.go index 834170a..95915aa 100644 --- a/internal/cli/auth_token.go +++ b/internal/cli/auth_token.go @@ -1,23 +1,46 @@ package cli import ( + "errors" "fmt" + "github.com/alyraffauf/tg/atproto" + "github.com/bluesky-social/indigo/atproto/atclient" "github.com/spf13/cobra" ) var authTokenCmd = &cobra.Command{ Use: "token", - Short: "Print the current OAuth access token", + Short: "Print the current access token", Args: cobra.NoArgs, RunE: func(cmd *cobra.Command, _ []string) error { - session, err := requireAuthSession(cmd.Context()) + ctx := cmd.Context() + session, err := auth.CurrentSession(ctx) + if err == nil { + token, _ := session.GetHostAccessData() + if token == "" { + return fmt.Errorf("current session has no access token") + } + fmt.Fprintln(cmd.OutOrStdout(), token) + return nil + } + if !errors.Is(err, atproto.ErrNotAuthenticated) { + return fmt.Errorf("resume OAuth session: %w", err) + } + client, _, err := auth.APIClient(ctx) if err != nil { - return err + if errors.Is(err, atproto.ErrNotAuthenticated) { + return fmt.Errorf("not logged in; run \"tg auth login\" first") + } + return fmt.Errorf("resume auth session: %w", err) + } + passwordAuth, ok := client.Auth.(*atclient.PasswordAuth) + if !ok { + return fmt.Errorf("not logged in; run \"tg auth login\" first") } - token, _ := session.GetHostAccessData() + token, _ := passwordAuth.GetTokens() if token == "" { - return fmt.Errorf("current OAuth session has no access token") + return fmt.Errorf("current session has no access token") } fmt.Fprintln(cmd.OutOrStdout(), token) return nil diff --git a/internal/cli/pr_create.go b/internal/cli/pr_create.go index f318e13..37ffef7 100644 --- a/internal/cli/pr_create.go +++ b/internal/cli/pr_create.go @@ -16,20 +16,21 @@ import ( const patchMimeType = "application/gzip" var ( - prCreateTitle string - prCreateBody string - prCreateBodyFile string - prCreateBase string - prCreateHead string - prCreateRepo string + prCreateTitle string + prCreateBody string + prCreateBodyFile string + prCreateBase string + prCreateHead string + prCreateRepo string + prCreateSourceRepo string ) var prCreateCmd = &cobra.Command{ Use: "create", Short: "Create a pull request from the current branch", Long: "Create a pull request by uploading a gzipped git patch and writing a sh.tangled.repo.pull record. " + - "The source and target repository are the same. By default, the current branch is the source and " + - "origin's default branch is the target. Use --repo to target a different Tangled repository.", + "By default, the current repository and branch are both the source and target repository, and origin's " + + "default branch is the target branch. Use --repo and --source-repo for a fork-based pull request.", Args: cobra.NoArgs, RunE: func(cmd *cobra.Command, args []string) error { ctx := cmd.Context() @@ -70,6 +71,20 @@ var prCreateCmd = &cobra.Command{ if !strings.HasPrefix(target.URI, "at://") { return fmt.Errorf("target repository %q has no strong at:// URI", repo) } + source := target + if prCreateSourceRepo != "" { + sourceHandle, sourceName, err := parseHandleRepo(prCreateSourceRepo) + if err != nil { + return err + } + source, err = resolveRepoRecord(ctx, sourceHandle, sourceName) + if err != nil { + return fmt.Errorf("resolve source repository: %w", err) + } + } + if source.Value.RepoDid == "" { + return fmt.Errorf("source repository has no repo DID") + } patch, err := gitutil.GeneratePatch(ctx, repoDir, base, head) if err != nil { @@ -81,12 +96,13 @@ var prCreateCmd = &cobra.Command{ } uri, err := createPullRecord(ctx, atClient, did, prCreateRecord{ - Title: prCreateTitle, - Body: body, - RepoDid: target.Value.RepoDid, - Base: base, - Head: head, - Patch: blob, + Title: prCreateTitle, + Body: body, + TargetRepoDid: target.Value.RepoDid, + SourceRepoDid: source.Value.RepoDid, + Base: base, + Head: head, + Patch: blob, }) if err != nil { return err @@ -105,16 +121,18 @@ func init() { prCreateCmd.Flags().StringVarP(&prCreateBase, "base", "B", "", "Target branch (default: origin's default branch)") prCreateCmd.Flags().StringVarP(&prCreateHead, "head", "H", "", "Source branch (default: current branch)") prCreateCmd.Flags().StringVarP(&prCreateRepo, "repo", "R", "", "Target repository as handle/repo") + prCreateCmd.Flags().StringVar(&prCreateSourceRepo, "source-repo", "", "Source repository as handle/repo (for fork-based pull requests)") prCreateCmd.MarkFlagRequired("title") } type prCreateRecord struct { - Title string - Body string - RepoDid string - Base string - Head string - Patch *atproto.Blob + Title string + Body string + TargetRepoDid string + SourceRepoDid string + Base string + Head string + Patch *atproto.Blob } // pullRecord is the sh.tangled.repo.pull lexicon shape used for record writes. @@ -129,15 +147,13 @@ type pullRecord struct { } type pullTarget struct { - Repo string `json:"repo"` - RepoDid string `json:"repoDid"` - Branch string `json:"branch"` + Repo string `json:"repo"` + Branch string `json:"branch"` } type pullSource struct { - Repo string `json:"repo"` - RepoDid string `json:"repoDid"` - Branch string `json:"branch"` + Repo string `json:"repo,omitempty"` + Branch string `json:"branch"` } type pullRound struct { @@ -175,35 +191,37 @@ func prTargetBranch(ctx context.Context, repoDir string) (string, error) { } func createPullRecord(ctx context.Context, atClient *atproto.ATProto, did string, input prCreateRecord) (string, error) { - now := time.Now().UTC().Format(time.RFC3339) - record := pullRecord{ + record := newPullRecord(input, time.Now().UTC()) + uri, _, err := atClient.PutRecord(ctx, atproto.PutRecordInput{ + Repo: did, + Collection: "sh.tangled.repo.pull", + Rkey: string(syntax.NewTIDNow(0)), + Record: record, + }) + if err != nil { + return "", fmt.Errorf("create pull request record: %w", err) + } + return uri, nil +} + +func newPullRecord(input prCreateRecord, createdAt time.Time) pullRecord { + now := createdAt.Format(time.RFC3339) + return pullRecord{ Type: "sh.tangled.repo.pull", Title: input.Title, Body: input.Body, CreatedAt: now, Target: pullTarget{ - Repo: input.RepoDid, - RepoDid: input.RepoDid, - Branch: input.Base, + Repo: input.TargetRepoDid, + Branch: input.Base, }, Source: pullSource{ - Repo: input.RepoDid, - RepoDid: input.RepoDid, - Branch: input.Head, + Repo: input.SourceRepoDid, + Branch: input.Head, }, Rounds: []pullRound{{ CreatedAt: now, PatchBlob: input.Patch, }}, } - uri, _, err := atClient.PutRecord(ctx, atproto.PutRecordInput{ - Repo: did, - Collection: "sh.tangled.repo.pull", - Rkey: string(syntax.NewTIDNow(0)), - Record: record, - }) - if err != nil { - return "", fmt.Errorf("create pull request record: %w", err) - } - return uri, nil } diff --git a/internal/cli/pr_create_test.go b/internal/cli/pr_create_test.go new file mode 100644 index 0000000..85389ec --- /dev/null +++ b/internal/cli/pr_create_test.go @@ -0,0 +1,26 @@ +package cli + +import ( + "testing" + "time" + + "github.com/alyraffauf/tg/atproto" +) + +func TestNewPullRecordUsesDistinctSourceAndTarget(t *testing.T) { + record := newPullRecord(prCreateRecord{ + Title: "Cross-repo change", + TargetRepoDid: "did:plc:upstream", + SourceRepoDid: "did:plc:fork", + Base: "main", + Head: "feature", + Patch: &atproto.Blob{}, + }, time.Date(2026, 7, 17, 0, 0, 0, 0, time.UTC)) + + if record.Target.Repo != "did:plc:upstream" { + t.Fatalf("unexpected target: %+v", record.Target) + } + if record.Source.Repo != "did:plc:fork" { + t.Fatalf("unexpected source: %+v", record.Source) + } +} diff --git a/internal/cli/repo_fork.go b/internal/cli/repo_fork.go index e73e251..c055d54 100644 --- a/internal/cli/repo_fork.go +++ b/internal/cli/repo_fork.go @@ -3,6 +3,7 @@ package cli import ( "context" "fmt" + "strings" "time" "github.com/alyraffauf/tg/atproto" @@ -42,7 +43,7 @@ var repoForkCmd = &cobra.Command{ repoDID, err := knot.New(source.Knot, token).CreateRepo(ctx, knot.CreateRepoInput{ Name: name, Rkey: name, - Source: source.URI, + Source: forkSourceURL(source.Knot, source.RepoDID), }) if err != nil { return err @@ -57,6 +58,7 @@ var repoForkCmd = &cobra.Command{ Knot: source.Knot, CreatedAt: time.Now().UTC().Format(time.RFC3339), RepoDid: repoDID, + Source: source.URI, }, }) if err != nil { @@ -85,8 +87,17 @@ func deleteFork(ctx context.Context, atClient *atproto.ATProto, knotHost, did, n } type forkSource struct { - URI string - Knot string + URI string + Knot string + RepoDID string +} + +func forkSourceURL(knotHost, repoDID string) string { + base := strings.TrimRight(knotHost, "/") + if !strings.HasPrefix(base, "http://") && !strings.HasPrefix(base, "https://") { + base = "https://" + base + } + return base + "/" + repoDID } func getForkSource(ctx context.Context, handle, name string) (forkSource, error) { @@ -102,10 +113,13 @@ func getForkSource(ctx context.Context, handle, name string) (forkSource, error) if repo.Value.Knot == "" { return forkSource{}, fmt.Errorf("source repository %s/%s has no knot", handle, name) } + if repo.Value.RepoDid == "" { + return forkSource{}, fmt.Errorf("source repository %s/%s has no repo DID", handle, name) + } if repo.URI != "" { uri = repo.URI } - return forkSource{URI: uri, Knot: repo.Value.Knot}, nil + return forkSource{URI: uri, Knot: repo.Value.Knot, RepoDID: repo.Value.RepoDid}, nil } type repoForkResult struct { diff --git a/internal/cli/repo_fork_test.go b/internal/cli/repo_fork_test.go new file mode 100644 index 0000000..14c6a68 --- /dev/null +++ b/internal/cli/repo_fork_test.go @@ -0,0 +1,24 @@ +package cli + +import "testing" + +func TestForkSourceURL(t *testing.T) { + tests := []struct { + name string + knot string + repoDID string + want string + }{ + {"bare host", "knot.gaze.systems", "did:plc:abc", "https://knot.gaze.systems/did:plc:abc"}, + {"https host", "https://knot.gaze.systems", "did:plc:abc", "https://knot.gaze.systems/did:plc:abc"}, + {"trailing slash", "https://knot.gaze.systems/", "did:plc:abc", "https://knot.gaze.systems/did:plc:abc"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := forkSourceURL(tt.knot, tt.repoDID); got != tt.want { + t.Fatalf("forkSourceURL() = %q, want %q", got, tt.want) + } + }) + } +} diff --git a/tangled/get_repo.go b/tangled/get_repo.go index 32cc51d..9f36848 100644 --- a/tangled/get_repo.go +++ b/tangled/get_repo.go @@ -17,6 +17,7 @@ type RepoRecord struct { Owner string `json:"owner,omitempty"` AddedAt string `json:"addedAt,omitempty"` RepoDid string `json:"repoDid,omitempty"` + Source string `json:"source,omitempty"` Spindle string `json:"spindle,omitempty"` Website string `json:"website,omitempty"` Labels []string `json:"labels,omitempty"`