mirror of
https://tangled.org/evan.jarrett.net/at-container-registry
synced 2026-08-31 05:07:09 +00:00
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>
397 lines
12 KiB
Go
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)
|
|
}
|
|
})
|
|
}
|
|
}
|