Files
at-container-registry/pkg/auth/oauth/client_test.go
T
Evan JarrettandClaude Fable 5 37bab324d7 fix OAuth refresh-token burn on client cancellation causing sign-outs
When a Docker client canceled a slow /auth/token request mid-refresh, the
token-refresh POST was aborted client-side but completed on the PDS, which
rotated the refresh token. The rotated token was never received or persisted,
so the next refresh replayed the consumed token, got invalid_grant, and the
session (OAuth + UI) was deleted, signing the user out everywhere.

- Detach refresh POSTs from the inbound request context via a per-session
  RoundTripper (WithoutCancel + 30s cap); once a refresh starts it completes
- Persist session updates (rotated tokens, DPoP nonces) on a detached context
- Gate session deletion on IsSessionInvalidError: cancellation, timeouts, and
  transport errors no longer delete sessions; genuine invalid_grant still does
- Add phase timing to /auth/token and per-DID lock wait warnings to attribute
  the ~14s pre-refresh stalls that push requests past Docker's deadline

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-02 13:38:45 -05:00

397 lines
12 KiB
Go

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},
// 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)
}
})
}
}