package oauth import ( "context" "errors" "fmt" "net/http" "net/http/httptest" "sync" "testing" "time" "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/atcrypto" "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/syntax" ) func TestNewClientApp(t *testing.T) { keyPath := t.TempDir() + "/oauth-key.bin" store := oauth.NewMemStore() baseURL := "http://localhost:5000" scopes := GetDefaultScopes("*") clientApp, err := NewClientApp(baseURL, store, scopes, keyPath, "AT Container Registry") if err != nil { t.Fatalf("NewClientApp() error = %v", err) } if clientApp == nil { t.Fatal("Expected non-nil clientApp") } if clientApp.Dir == nil { t.Error("Expected directory to be set") } } func TestNewClientAppWithCustomScopes(t *testing.T) { keyPath := t.TempDir() + "/oauth-key.bin" store := oauth.NewMemStore() baseURL := "http://localhost:5000" scopes := []string{"atproto", "custom:scope"} clientApp, err := NewClientApp(baseURL, store, scopes, keyPath, "AT Container Registry") if err != nil { t.Fatalf("NewClientApp() error = %v", err) } if clientApp == nil { t.Fatal("Expected non-nil clientApp") } // Verify clientApp was created successfully // (Note: indigo's oauth.ClientApp doesn't expose scopes directly, // but we can verify it was created without error) if clientApp.Dir == nil { t.Error("Expected directory to be set") } } func TestScopesMatch(t *testing.T) { tests := []struct { name string stored []string desired []string expected bool }{ { name: "exact match", stored: []string{"atproto", "blob:image/png"}, desired: []string{"atproto", "blob:image/png"}, expected: true, }, { name: "different order", stored: []string{"blob:image/png", "atproto"}, desired: []string{"atproto", "blob:image/png"}, expected: true, }, { name: "missing scope in stored", stored: []string{"atproto"}, desired: []string{"atproto", "blob:image/png"}, expected: false, }, { name: "extra scope in stored", stored: []string{"atproto", "blob:image/png", "extra"}, desired: []string{"atproto", "blob:image/png"}, expected: false, }, { name: "both empty", stored: []string{}, desired: []string{}, expected: true, }, { name: "nil vs empty", stored: nil, desired: []string{}, expected: true, }, { name: "completely different", stored: []string{"foo", "bar"}, desired: []string{"baz", "qux"}, expected: false, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { result := ScopesMatch(tt.stored, tt.desired) if result != tt.expected { t.Errorf("ScopesMatch(%v, %v) = %v, want %v", tt.stored, tt.desired, result, tt.expected) } }) } } // ---------------------------------------------------------------------------- // Session Management (Refresher) Tests // ---------------------------------------------------------------------------- func TestNewRefresher(t *testing.T) { store := oauth.NewMemStore() scopes := GetDefaultScopes("*") clientApp, err := NewClientApp("http://localhost:5000", store, scopes, "", "AT Container Registry") if err != nil { t.Fatalf("NewClientApp() error = %v", err) } refresher := NewRefresher(clientApp) if refresher == nil { t.Fatal("Expected non-nil refresher") } if refresher.clientApp == nil { t.Error("Expected clientApp to be set") } } func TestRefresher_SetUISessionStore(t *testing.T) { store := oauth.NewMemStore() scopes := GetDefaultScopes("*") clientApp, err := NewClientApp("http://localhost:5000", store, scopes, "", "AT Container Registry") if err != nil { t.Fatalf("NewClientApp() error = %v", err) } refresher := NewRefresher(clientApp) // Test that SetUISessionStore doesn't panic with nil // Full mock implementation requires implementing the interface refresher.SetUISessionStore(nil) // Verify nil is accepted if refresher.uiSessionStore != nil { t.Error("Expected UI session store to be nil after setting nil") } } // ---------------------------------------------------------------------------- // Refresh-cancellation regression tests // ---------------------------------------------------------------------------- // fakeAuthStore is a ClientAuthStore whose SaveSession honors context // cancellation, so a test fails if session persistence runs on a canceled // context. It also implements GetLatestSessionForDID (the sessionGetter // extension the Refresher requires). type fakeAuthStore struct { mu sync.Mutex sessions map[string]oauth.ClientSessionData // keyed by DID deleted []string } func newFakeAuthStore() *fakeAuthStore { return &fakeAuthStore{sessions: make(map[string]oauth.ClientSessionData)} } func (s *fakeAuthStore) GetSession(ctx context.Context, did syntax.DID, sessionID string) (*oauth.ClientSessionData, error) { s.mu.Lock() defer s.mu.Unlock() sess, ok := s.sessions[did.String()] if !ok { return nil, fmt.Errorf("session not found") } return &sess, nil } func (s *fakeAuthStore) SaveSession(ctx context.Context, sess oauth.ClientSessionData) error { if err := ctx.Err(); err != nil { return err } s.mu.Lock() defer s.mu.Unlock() s.sessions[sess.AccountDID.String()] = sess return nil } func (s *fakeAuthStore) DeleteSession(ctx context.Context, did syntax.DID, sessionID string) error { s.mu.Lock() defer s.mu.Unlock() delete(s.sessions, did.String()) s.deleted = append(s.deleted, did.String()) return nil } func (s *fakeAuthStore) GetAuthRequestInfo(ctx context.Context, state string) (*oauth.AuthRequestData, error) { return nil, fmt.Errorf("not implemented") } func (s *fakeAuthStore) SaveAuthRequestInfo(ctx context.Context, info oauth.AuthRequestData) error { return nil } func (s *fakeAuthStore) DeleteAuthRequestInfo(ctx context.Context, state string) error { return nil } func (s *fakeAuthStore) GetLatestSessionForDID(ctx context.Context, did string) (*oauth.ClientSessionData, string, error) { s.mu.Lock() defer s.mu.Unlock() sess, ok := s.sessions[did] if !ok { return nil, "", fmt.Errorf("no session for DID") } return &sess, sess.SessionID, nil } func (s *fakeAuthStore) refreshToken(did string) string { s.mu.Lock() defer s.mu.Unlock() return s.sessions[did].RefreshToken } func (s *fakeAuthStore) deletedDIDs() []string { s.mu.Lock() defer s.mu.Unlock() return append([]string(nil), s.deleted...) } // spyUISessionStore records DeleteByDID calls. type spyUISessionStore struct { mu sync.Mutex deleted []string } func (s *spyUISessionStore) Create(did, handle, pdsEndpoint string, duration time.Duration) (string, error) { return "", fmt.Errorf("not implemented") } func (s *spyUISessionStore) DeleteByDID(did string) { s.mu.Lock() defer s.mu.Unlock() s.deleted = append(s.deleted, did) } // seedSession stores a session for did pointing at the given resource server // and token endpoint, with a freshly generated DPoP key. func seedSession(t *testing.T, store *fakeAuthStore, did, hostURL, tokenEndpoint string) { t.Helper() key, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } parsedDID, err := syntax.ParseDID(did) if err != nil { t.Fatal(err) } store.sessions[did] = oauth.ClientSessionData{ AccountDID: parsedDID, SessionID: "test-session", HostURL: hostURL, AuthServerURL: hostURL, AuthServerTokenEndpoint: tokenEndpoint, Scopes: []string{"atproto"}, AccessToken: "old-access", RefreshToken: "old-refresh", DPoPPrivateKeyMultibase: key.Multibase(), } } func newTestRefresher(t *testing.T, store oauth.ClientAuthStore) *Refresher { t.Helper() clientApp, err := NewClientApp("http://localhost:5000", store, []string{"atproto"}, "", "test") if err != nil { t.Fatal(err) } return NewRefresher(clientApp) } // TestDoWithSession_RefreshSurvivesRequestCancellation reproduces the // production incident: the caller's context is canceled while the token // refresh is in flight (after the auth server has already rotated the // refresh token). The rotated token MUST be persisted and the session MUST // NOT be deleted, or the next refresh fails with invalid_grant "Refresh // token replayed" and the user is signed out. func TestDoWithSession_RefreshSurvivesRequestCancellation(t *testing.T) { const did = "did:web:refresh-cancel.example.com" ctx, cancel := context.WithCancel(context.Background()) defer cancel() mux := http.NewServeMux() srv := httptest.NewServer(mux) defer srv.Close() // Resource endpoint: reject the stale access token so DoWithAuth triggers // a refresh; accept the rotated one. mux.HandleFunc("/xrpc/test.resource", func(w http.ResponseWriter, r *http.Request) { if r.Header.Get("Authorization") == "DPoP new-access" { w.WriteHeader(http.StatusOK) return } w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token", error_description="expired"`) w.WriteHeader(http.StatusUnauthorized) }) // Token endpoint: simulate the Docker client hanging up mid-refresh by // canceling the caller's context BEFORE responding, then return rotated // tokens (the auth server has already committed the rotation by then). mux.HandleFunc("/oauth/token", func(w http.ResponseWriter, r *http.Request) { cancel() fmt.Fprintf(w, `{"sub":%q,"access_token":"new-access","refresh_token":"new-refresh"}`, did) }) store := newFakeAuthStore() seedSession(t, store, did, srv.URL, srv.URL+"/oauth/token") refresher := newTestRefresher(t, store) uiStore := &spyUISessionStore{} refresher.SetUISessionStore(uiStore) err := refresher.DoWithSession(ctx, did, func(session *oauth.ClientSession) error { req, err := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL+"/xrpc/test.resource", nil) if err != nil { return err } resp, err := session.DoWithAuth(session.Client, req, syntax.NSID("com.atproto.server.getServiceAuth")) if err != nil { return err } defer resp.Body.Close() return nil }) // The overall operation may fail (the post-refresh retry of the resource // request runs on the canceled inbound context) — that is fine, Docker // retries. What must hold is that the rotated refresh token was saved and // the session survived. if got := store.refreshToken(did); got != "new-refresh" { t.Errorf("rotated refresh token not persisted: got %q, want %q (op err: %v)", got, "new-refresh", err) } if deleted := store.deletedDIDs(); len(deleted) != 0 { t.Errorf("session was deleted: %v", deleted) } if len(uiStore.deleted) != 0 { t.Errorf("UI session was invalidated: %v", uiStore.deleted) } } func TestIsSessionInvalidError(t *testing.T) { tests := []struct { name string err error want bool }{ {"nil", nil, false}, {"plain canceled", context.Canceled, false}, {"wrapped canceled", fmt.Errorf("token refresh failed: %w", context.Canceled), false}, {"wrapped deadline", fmt.Errorf("fetch: %w", context.DeadlineExceeded), false}, // Even if the message mentions an auth string, cancellation wins. {"canceled with auth-ish text", fmt.Errorf("invalid_grant: %w", context.Canceled), false}, {"api error 401", &atclient.APIError{StatusCode: 401}, true}, {"api error InvalidGrant", &atclient.APIError{StatusCode: 400, Name: "InvalidGrant"}, true}, {"api error InvalidToken", &atclient.APIError{StatusCode: 400, Name: "InvalidToken"}, true}, {"api error 500", &atclient.APIError{StatusCode: 500, Name: "InternalServerError"}, false}, // ExpiredToken means "refresh me", not "revoked". Treating it as a dead // session signs the user out of every UI session over an ordinary // access-token expiry that a refresh would have fixed. {"api error ExpiredToken is refreshable, not dead", &atclient.APIError{StatusCode: 400, Name: "ExpiredToken"}, false}, // Transient upstream failures must never evict: these are the shapes the // service-token path now wraps as APIErrors. {"api error 502", &atclient.APIError{StatusCode: 502, Name: ""}, false}, {"api error 429", &atclient.APIError{StatusCode: 429, Name: ""}, false}, {"api error 500 html body", &atclient.APIError{StatusCode: 500, Name: "", Message: "bad gateway"}, false}, // A revoked session reported as 401 with an atproto name — the case the // service-token path was previously flattening into an unmatchable string. {"api error 401 InvalidToken", &atclient.APIError{StatusCode: 401, Name: "InvalidToken"}, true}, // The refresh-replay failure arrives as a plain wrapped string from indigo. {"plain invalid_grant string", errors.New("failed to refresh OAuth tokens: token refresh failed (HTTP 400): invalid_grant"), true}, {"plain invalid_token string", errors.New("auth server request failed (HTTP 401): invalid_token"), true}, {"connection refused", errors.New(`Post "https://pds.example.com/oauth/token": dial tcp: connection refused`), false}, {"generic 500", errors.New("token refresh failed (HTTP 500): server exploded"), false}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { if got := IsSessionInvalidError(tt.err); got != tt.want { t.Errorf("IsSessionInvalidError(%v) = %v, want %v", tt.err, got, tt.want) } }) } }