package oauth import ( "context" "crypto/rand" "encoding/base64" "encoding/json" "errors" "fmt" "net/http" "net/http/httptest" "net/url" "strings" "sync" "testing" "time" "github.com/bluesky-social/indigo/atproto/atcrypto" "github.com/bluesky-social/indigo/atproto/auth" indigooauth "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/identity" "github.com/bluesky-social/indigo/atproto/syntax" "tangled.org/core/log" "tangled.org/core/migrator/config" "tangled.org/core/migrator/db" ) type grantStub struct { mu sync.Mutex server *httptest.Server families int refreshes int usedRefresh map[string]bool latestAccess string mintCalls []string tokenStatus int tokenBody string invalidGrantDescription string userDID syntax.DID userKey atcrypto.PrivateKey } func newGrantStub(t *testing.T, userDID syntax.DID) *grantStub { t.Helper() userKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } g := &grantStub{ usedRefresh: make(map[string]bool), userDID: userDID, userKey: userKey, } g.server = httptest.NewTLSServer(http.HandlerFunc(g.serve)) t.Cleanup(g.server.Close) return g } func (g *grantStub) serve(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { case "/oauth/token": g.serveToken(w, r) case "/xrpc/com.atproto.server.getServiceAuth": g.serveServiceAuth(w, r) default: http.NotFound(w, r) } } func (g *grantStub) serveToken(w http.ResponseWriter, r *http.Request) { if err := r.ParseForm(); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } g.mu.Lock() defer g.mu.Unlock() if g.tokenStatus != 0 { w.Header().Set("Content-Type", "application/json") w.WriteHeader(g.tokenStatus) _, _ = w.Write([]byte(g.tokenBody)) return } if r.Form.Get("grant_type") == "authorization_code" { g.families++ g.latestAccess = fmt.Sprintf("pds-oauth-family-%d", g.families) _ = json.NewEncoder(w).Encode(indigooauth.TokenResponse{ Subject: g.userDID.String(), Scope: strings.Join(scopes, " "), AccessToken: g.latestAccess, RefreshToken: fmt.Sprintf("refresh-family-%d", g.families), }) return } refresh := r.Form.Get("refresh_token") if g.usedRefresh[refresh] { description := g.invalidGrantDescription if description == "" { description = "Refresh token replayed" } w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusBadRequest) _, _ = fmt.Fprintf(w, `{"error":"invalid_grant","error_description":%q}`, description) return } g.usedRefresh[refresh] = true g.refreshes++ g.latestAccess = fmt.Sprintf("pds-oauth-refreshed-%d", g.refreshes) _ = json.NewEncoder(w).Encode(indigooauth.TokenResponse{ Subject: g.userDID.String(), Scope: strings.Join(scopes, " "), AccessToken: g.latestAccess, RefreshToken: fmt.Sprintf("refresh-next-%d", g.refreshes), }) } func (g *grantStub) serveServiceAuth(w http.ResponseWriter, r *http.Request) { g.mu.Lock() if bearer := strings.TrimPrefix(r.Header.Get("Authorization"), "DPoP "); bearer != g.latestAccess { g.mu.Unlock() w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token", error_description="The access token expired"`) w.WriteHeader(http.StatusUnauthorized) return } g.mintCalls = append(g.mintCalls, r.URL.Query().Get("lxm")) g.mu.Unlock() lxm, err := syntax.ParseNSID(r.URL.Query().Get("lxm")) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } aud := r.URL.Query().Get("aud") token, err := auth.SignServiceAuth(g.userDID, aud, time.Minute, &lxm, g.userKey) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } _ = json.NewEncoder(w).Encode(map[string]string{"token": token}) } func clientWithGrantStub(t *testing.T, g *grantStub, ownerDid syntax.DID) *Client { t.Helper() store, database := testStore(t) rawKey := make([]byte, 32) _, _ = rand.Read(rawKey) serviceKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } cfg := &config.Config{ Hostname: "migrator.example.com", PrivateKey: serviceKey.Multibase(), MasterKey: base64.StdEncoding.EncodeToString(rawKey), WorkDir: t.TempDir(), } if err := cfg.Validate(); err != nil { t.Fatal(err) } client, err := NewClient(cfg, database, identity.NewMockDirectory(), log.New("oauth-session-test")) if err != nil { t.Fatal(err) } client.app.Client = g.server.Client() client.http = g.server.Client() client.store = store client.app.Store = store dpopKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } if err := store.SaveSession(context.Background(), indigooauth.ClientSessionData{ AccountDID: ownerDid, SessionID: "session", HostURL: g.server.URL, AuthServerURL: g.server.URL, AuthServerTokenEndpoint: g.server.URL + "/oauth/token", Scopes: scopes, AccessToken: "pds-oauth-old", RefreshToken: "refresh-old", DPoPPrivateKeyMultibase: dpopKey.Multibase(), }); err != nil { t.Fatal(err) } return client } func TestConcurrentJobsForOneOwnerRefreshExactlyOnce(t *testing.T) { ownerDid := syntax.DID("did:plc:alice") g := newGrantStub(t, ownerDid) client := clientWithGrantStub(t, g, ownerDid) const callers = 4 errs := make([]error, callers) var wg sync.WaitGroup for i := range callers { wg.Add(1) go func(i int) { defer wg.Done() _, errs[i] = client.userKnotToken(context.Background(), ownerDid.String(), "did:web:knot.example", "com.example.lxm") }(i) } wg.Wait() for i, err := range errs { if err != nil { t.Fatalf("caller %d failed: %v", i, err) } } g.mu.Lock() defer g.mu.Unlock() if g.refreshes != 1 { t.Fatalf("token refreshes = %d, want exactly one for the whole token family", g.refreshes) } for _, lxm := range g.mintCalls { if lxm != "com.example.lxm" { t.Fatalf("unexpected mint call lxm %q", lxm) } } if len(g.mintCalls) != callers { t.Fatalf("service auth calls = %d, want %d", len(g.mintCalls), callers) } } func TestARefusedGrantIsTerminal(t *testing.T) { for _, description := range []string{"Refresh token replayed", "Invalid refresh token"} { t.Run(description, func(t *testing.T) { ownerDid := syntax.DID("did:plc:alice") g := newGrantStub(t, ownerDid) g.invalidGrantDescription = description g.tokenStatus = http.StatusBadRequest g.tokenBody = fmt.Sprintf(`{"error":"invalid_grant","error_description":%q}`, description) client := clientWithGrantStub(t, g, ownerDid) _, err := client.userKnotToken(context.Background(), ownerDid.String(), "did:web:knot.example", "com.example.lxm") if !errors.Is(err, ErrGrantRequired) { t.Fatalf("err = %v, want ErrGrantRequired", err) } if _, err := client.store.SessionID(context.Background(), ownerDid.String()); !errors.Is(err, db.ErrNotFound) { t.Fatalf("session row after refusal = %v, want gone", err) } if client.HasSession(context.Background(), ownerDid.String()) { t.Fatal("a refused grant still reports a usable session") } }) } } func TestAnAuthServerOutageStaysRetryable(t *testing.T) { ownerDid := syntax.DID("did:plc:alice") g := newGrantStub(t, ownerDid) g.tokenStatus = http.StatusServiceUnavailable g.tokenBody = `{"error":"temporarily_unavailable"}` client := clientWithGrantStub(t, g, ownerDid) _, err := client.userKnotToken(context.Background(), ownerDid.String(), "did:web:knot.example", "com.example.lxm") if err == nil { t.Fatal("expected an error from the auth server outage") } if errors.Is(err, ErrGrantRequired) { t.Fatalf("a 5xx outage became a refused grant: %v", err) } } func TestARegrantReplacesADeadCachedSession(t *testing.T) { ownerDid := syntax.DID("did:plc:alice") g := newGrantStub(t, ownerDid) client := clientWithGrantStub(t, g, ownerDid) dir := identity.NewMockDirectory() dir.Insert(identity.Identity{ DID: ownerDid, Services: map[string]identity.ServiceEndpoint{ "atproto_pds": {Type: "AtprotoPersonalDataServer", URL: g.server.URL}, }, }) client.app.Dir = dir replay := func(refresh string) { t.Helper() form := url.Values{} form.Set("grant_type", "refresh_token") form.Set("refresh_token", refresh) form.Set("client_id", client.app.Config.ClientID) resp, err := g.server.Client().PostForm(g.server.URL+"/oauth/token", form) if err != nil { t.Fatal(err) } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("replaying %q: status %d", refresh, resp.StatusCode) } } replay("refresh-old") stale, err := client.session(context.Background(), ownerDid.String()) if err != nil { t.Fatal(err) } dpopKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } info := indigooauth.AuthRequestData{ State: "session-2", AuthServerURL: g.server.URL, AccountDID: &ownerDid, Scopes: scopes, RequestURI: "request-uri", AuthServerTokenEndpoint: g.server.URL + "/oauth/token", PKCEVerifier: "verifier", DPoPPrivateKeyMultibase: dpopKey.Multibase(), } if err := client.store.SaveAuthRequestInfo(context.Background(), info); err != nil { t.Fatal(err) } req := httptest.NewRequest(http.MethodGet, "/oauth/callback?state=session-2&iss="+url.QueryEscape(g.server.URL)+"&code=xyz", nil) rec := httptest.NewRecorder() client.Routes().ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("re-grant callback = %d %s, want 303", rec.Code, rec.Body.String()) } if strings.Contains(rec.Header().Get("Location"), "oauth_error") { t.Fatalf("re-grant failed: %q", rec.Header().Get("Location")) } if id, err := client.store.SessionID(context.Background(), ownerDid.String()); err != nil || id != "session-2" { t.Fatalf("session id after re-grant = %q err=%v, want session-2", id, err) } // simulate a resume race re-pinning the stale session after callback eviction client.sessions.Store(ownerDid.String(), stale) g.mu.Lock() afterRegrant := g.refreshes g.mu.Unlock() if err := client.UsableGrant(context.Background(), ownerDid.String()); err != nil { t.Fatalf("probe after re-grant: %v", err) } if _, err := client.userKnotToken(context.Background(), ownerDid.String(), "did:web:knot.example", "com.example.lxm"); err != nil { t.Fatalf("mint after re-grant: %v", err) } g.mu.Lock() defer g.mu.Unlock() if g.refreshes != afterRegrant { t.Fatalf("token refreshes after re-grant = %d, want %d (the fresh family must not need one)", g.refreshes, afterRegrant) } }