From dab96ca1f5c1a1d87beae567fd68750c9d31ecc2 Mon Sep 17 00:00:00 2001 From: Aly Raffauf Date: Wed, 22 Jul 2026 00:51:05 -0400 Subject: [PATCH] reuse custom httpClient across stack --- atproto/auth.go | 26 +++++++++++++++++-- atproto/auth_test.go | 49 +++++++++++++++++++++++++++++++++++ atproto/keyring_store_test.go | 6 +++-- internal/app/app.go | 14 ++++++---- internal/app/dependencies.go | 16 +++++++----- knot/client.go | 10 +++++-- 6 files changed, 104 insertions(+), 17 deletions(-) diff --git a/atproto/auth.go b/atproto/auth.go index 32db002..73fe141 100644 --- a/atproto/auth.go +++ b/atproto/auth.go @@ -83,6 +83,7 @@ var DefaultScopes = []string{ type AuthManager struct { app *oauth.ClientApp store *KeyringStore + client *http.Client selector string pendingIdentifier string } @@ -108,12 +109,21 @@ func (m *AuthManager) activeAccount() (Account, error) { } func NewAuthManager(callbackURL string) *AuthManager { + return NewAuthManagerWithClient(callbackURL, http.DefaultClient) +} + +// NewAuthManagerWithClient creates an AuthManager using httpClient for OAuth +// and authenticated API requests. +func NewAuthManagerWithClient(callbackURL string, httpClient *http.Client) *AuthManager { config := oauth.NewLocalhostConfig(callbackURL, DefaultScopes) config.UserAgent = "tg" store := NewKeyringStore() + app := oauth.NewClientApp(&config, store) + app.Client = httpClient return &AuthManager{ - app: oauth.NewClientApp(&config, store), - store: store, + app: app, + store: store, + client: httpClient, } } @@ -125,6 +135,8 @@ func (m *AuthManager) LoginWithPassword(ctx context.Context, identifier, passwor if err != nil { return err } + ctx, cancel := m.requestContext(ctx) + defer cancel() persist := func(_ context.Context, data atclient.PasswordSessionData) { _ = m.store.SavePasswordSession(context.Background(), data) } @@ -136,6 +148,7 @@ func (m *AuthManager) LoginWithPassword(ctx context.Context, identifier, passwor if !ok { return errors.New("password login returned an unexpected auth type") } + client.Client = m.client if err := m.store.SavePasswordSession(ctx, passwordAuth.Session); err != nil { return err } @@ -265,9 +278,17 @@ func (m *AuthManager) APIClient(ctx context.Context) (*atclient.APIClient, synta _ = m.store.SavePasswordSession(context.Background(), data) } client := atclient.ResumePasswordSession(*passwordSession, persist) + client.Client = m.client return client, passwordSession.AccountDID, nil } +func (m *AuthManager) requestContext(ctx context.Context) (context.Context, context.CancelFunc) { + if m.client == nil || m.client.Timeout <= 0 { + return ctx, func() {} + } + return context.WithTimeout(ctx, m.client.Timeout) +} + // SessionStatus probes the server to verify the active session. The probe // refreshes the access token if needed. Returns ErrNotAuthenticated when // there is no active account. @@ -318,6 +339,7 @@ func (m *AuthManager) Logout(ctx context.Context) error { } if passwordSession != nil { client := atclient.ResumePasswordSession(*passwordSession, nil) + client.Client = m.client if passwordAuth, ok := client.Auth.(*atclient.PasswordAuth); ok { _ = passwordAuth.Logout(ctx, client.Client) } diff --git a/atproto/auth_test.go b/atproto/auth_test.go index 108f5ad..acdedbb 100644 --- a/atproto/auth_test.go +++ b/atproto/auth_test.go @@ -7,6 +7,7 @@ import ( "net/http/httptest" "strings" "testing" + "time" "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/syntax" @@ -135,6 +136,22 @@ func TestAPIClientReturnsPasswordClient(t *testing.T) { } } +func TestRequestContextUsesClientTimeout(t *testing.T) { + manager := newAuthManagerForTest("http://127.0.0.1:8095/callback", testKeyringStore(newFakeKeyring())) + manager.client = &http.Client{Timeout: 30 * time.Second} + + ctx, cancel := manager.requestContext(context.Background()) + defer cancel() + + deadline, ok := ctx.Deadline() + if !ok { + t.Fatal("request context has no deadline") + } + if remaining := time.Until(deadline); remaining <= 0 || remaining > 30*time.Second { + t.Errorf("request context deadline is %v away, want within 30 seconds", remaining) + } +} + func TestAPIClientNotAuthenticated(t *testing.T) { manager := newAuthManagerForTest("http://127.0.0.1:8095/callback", testKeyringStore(newFakeKeyring())) _, _, err := manager.APIClient(context.Background()) @@ -181,6 +198,38 @@ func TestPasswordLogoutRevokesAndRemovesSession(t *testing.T) { } } +func TestPasswordLogoutUsesConfiguredClient(t *testing.T) { + store := testKeyringStore(newFakeKeyring()) + manager := newAuthManagerForTest("http://127.0.0.1:8095/callback", store) + called := false + manager.client = &http.Client{Transport: roundTripFunc(func(request *http.Request) (*http.Response, error) { + called = true + return &http.Response{ + StatusCode: http.StatusOK, + Body: http.NoBody, + Header: make(http.Header), + }, nil + })} + 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) + } + if err := manager.Logout(ctx); err != nil { + t.Fatalf("Logout: %v", err) + } + if !called { + t.Error("configured HTTP client was not used") + } +} + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { + return f(request) +} + func TestPasswordLogoutPreservesOtherAccount(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) diff --git a/atproto/keyring_store_test.go b/atproto/keyring_store_test.go index 81a6d43..7da3963 100644 --- a/atproto/keyring_store_test.go +++ b/atproto/keyring_store_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "net/http" "reflect" "strings" "sync" @@ -55,8 +56,9 @@ func newAuthManagerForTest(callbackURL string, store *KeyringStore) *AuthManager config := oauth.NewLocalhostConfig(callbackURL, DefaultScopes) config.UserAgent = "tg" return &AuthManager{ - app: oauth.NewClientApp(&config, store), - store: store, + app: oauth.NewClientApp(&config, store), + store: store, + client: http.DefaultClient, } } diff --git a/internal/app/app.go b/internal/app/app.go index e888ede..0b549cd 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -6,6 +6,7 @@ import ( "log/slog" "net/http" "os" + "time" "github.com/alyraffauf/tg/atproto" "github.com/alyraffauf/tg/internal/gitutil" @@ -28,6 +29,8 @@ type Service struct { // DefaultKnot is used when repository creation does not specify a knot. const DefaultKnot = knot.DefaultKnot +const defaultHTTPTimeout = 30 * time.Second + // New returns a Service with production defaults: the default atproto // identity directory, the given appview host, and an AuthManager using // oauthCallbackURL for localhost OAuth redirects. @@ -38,19 +41,20 @@ func New(appviewHost, oauthCallbackURL string) *Service { // NewWithStreams creates production dependencies with configurable command // output streams. func NewWithStreams(appviewHost, oauthCallbackURL string, stdout, stderr io.Writer) *Service { + httpClient := &http.Client{Timeout: defaultHTTPTimeout} resolver := &atproto.Resolver{Directory: identity.DefaultDirectory()} - auth := atproto.NewAuthManager(oauthCallbackURL) + auth := atproto.NewAuthManagerWithClient(oauthCallbackURL, httpClient) return &Service{ resolver: resolver, appview: &tangled.Tangled{ - Client: &atclient.APIClient{Host: appviewHost}, + Client: &atclient.APIClient{Client: httpClient, Host: appviewHost}, Logger: slog.Default(), }, - sessions: productionSessions{auth: auth, resolver: resolver}, + sessions: productionSessions{auth: auth, resolver: resolver, httpClient: httpClient}, auth: auth, git: gitutil.NewClient(stdout, stderr), - knot: productionKnotFactory{}, - httpClient: http.DefaultClient, + knot: productionKnotFactory{httpClient: httpClient}, + httpClient: httpClient, } } diff --git a/internal/app/dependencies.go b/internal/app/dependencies.go index a53162a..fd91ed1 100644 --- a/internal/app/dependencies.go +++ b/internal/app/dependencies.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "net/http" "github.com/alyraffauf/tg/atproto" "github.com/alyraffauf/tg/internal/gitutil" @@ -63,8 +64,9 @@ type knotClientFactory interface { } type productionSessions struct { - auth *atproto.AuthManager - resolver identityResolver + auth *atproto.AuthManager + resolver identityResolver + httpClient *http.Client } func (s productionSessions) AuthenticatedPDS(ctx context.Context) (pdsClient, string, error) { @@ -87,7 +89,7 @@ func (s productionSessions) PublicPDS(ctx context.Context, handle string) (pdsCl if err != nil { return nil, "", fmt.Errorf("resolve PDS for %q: %w", handle, err) } - return &atproto.ATProto{Client: &atclient.APIClient{Host: pdsURL}}, ident.DID.String(), nil + return &atproto.ATProto{Client: &atclient.APIClient{Client: s.httpClient, Host: pdsURL}}, ident.DID.String(), nil } func (s productionSessions) APIClient(ctx context.Context) (*atclient.APIClient, error) { @@ -101,10 +103,12 @@ func (s productionSessions) APIClient(ctx context.Context) (*atclient.APIClient, return client, nil } -type productionKnotFactory struct{} +type productionKnotFactory struct { + httpClient *http.Client +} -func (productionKnotFactory) New(host, token string) knotClient { - return knot.New(host, token) +func (f productionKnotFactory) New(host, token string) knotClient { + return knot.NewWithClient(host, token, f.httpClient) } func isNotAuthenticated(err error) bool { diff --git a/knot/client.go b/knot/client.go index ba700ca..d179da9 100644 --- a/knot/client.go +++ b/knot/client.go @@ -18,10 +18,16 @@ type Client struct { // New returns a Client for host, authenticated with a service-auth token. func New(host, token string) *Client { + return NewWithClient(host, token, http.DefaultClient) +} + +// NewWithClient returns a Client using httpClient for requests. +func NewWithClient(host, token string, httpClient *http.Client) *Client { return &Client{ APIClient: &atclient.APIClient{ - Host: "https://" + host, - Auth: bearerAuth(token), + Client: httpClient, + Host: "https://" + host, + Auth: bearerAuth(token), }, } } -- 2.51.2