package api import ( "context" cryptorand "crypto/rand" "encoding/base64" "encoding/json" "fmt" "net/http" "net/http/httptest" "path/filepath" "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/api/org_tangled" "tangled.org/core/log" "tangled.org/core/migrator/config" "tangled.org/core/migrator/db" "tangled.org/core/migrator/oauth" "tangled.org/core/xrpc/serviceauth" ) type grantStub struct { mu sync.Mutex server *httptest.Server dead bool refreshes int requests int usedRefresh map[string]bool latestAccess string mintCalls []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: map[string]bool{}, userDID: userDID, userKey: userKey} g.server = httptest.NewServer(http.HandlerFunc(g.serve)) t.Cleanup(g.server.Close) return g } func (g *grantStub) serve(w http.ResponseWriter, r *http.Request) { g.mu.Lock() g.requests++ g.mu.Unlock() switch r.URL.Path { case "/oauth/token": g.token(w, r) case "/xrpc/com.atproto.server.getServiceAuth": g.serviceAuth(w, r) default: http.NotFound(w, r) } } func (g *grantStub) token(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.dead || g.usedRefresh[r.Form.Get("refresh_token")] { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusBadRequest) _, _ = w.Write([]byte(`{"error":"invalid_grant","error_description":"Refresh token replayed"}`)) return } g.usedRefresh[r.Form.Get("refresh_token")] = true g.refreshes++ g.latestAccess = fmt.Sprintf("pds-oauth-refreshed-%d", g.refreshes) _ = json.NewEncoder(w).Encode(indigooauth.TokenResponse{ Subject: g.userDID.String(), Scope: "atproto", AccessToken: g.latestAccess, RefreshToken: fmt.Sprintf("refresh-next-%d", g.refreshes), }) } func (g *grantStub) serviceAuth(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")+" aud="+r.URL.Query().Get("aud")) g.mu.Unlock() lxm, err := syntax.ParseNSID(r.URL.Query().Get("lxm")) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } token, err := auth.SignServiceAuth(g.userDID, r.URL.Query().Get("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 serverWithGrants(t *testing.T, g *grantStub) (*testServer, string) { t.Helper() // the stub pds dials on loopback, which the oauth client's ssrf guard refuses t.Setenv("ATPROTO_OAUTH_DEV", "1") dir := t.TempDir() database, err := db.Make(context.Background(), filepath.Join(dir, "api_grant.db")) if err != nil { t.Fatal(err) } t.Cleanup(func() { database.Close() }) rawKey := make([]byte, 32) _, _ = cryptorand.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: dir, } if err := cfg.Validate(); err != nil { t.Fatal(err) } oauthClient, err := oauth.NewClient(cfg, database, identity.NewMockDirectory(), log.New("api-grant-test")) if err != nil { t.Fatal(err) } store := oauth.NewStore(database, cfg.ParsedMasterKey) dpopKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } owner := syntax.DID(alice) if err := store.SaveSession(context.Background(), indigooauth.ClientSessionData{ AccountDID: owner, SessionID: "session", HostURL: g.server.URL, AuthServerURL: g.server.URL, AuthServerTokenEndpoint: g.server.URL + "/oauth/token", Scopes: []string{ "atproto", "rpc:sh.tangled.repo.create?aud=*", "rpc:sh.tangled.repo.describeRepo?aud=*", "repo:sh.tangled.repo?action=create", }, AccessToken: "pds-oauth-old", RefreshToken: "refresh-old", DPoPPrivateKeyMultibase: dpopKey.Multibase(), }); err != nil { t.Fatal(err) } logger := log.New("test") directory := identity.NewMockDirectory() priv, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } pub, err := priv.PublicKey() if err != nil { t.Fatal(err) } directory.Insert(identity.Identity{ DID: owner, Keys: map[string]identity.VerificationMethod{ "atproto": {Type: "Multikey", PublicKeyMultibase: pub.Multibase()}, }, }) serviceAuth := serviceauth.NewServiceAuth(logger, directory, cfg.ServiceDid.String()) return &testServer{ server: NewServer(database, cfg, nil, serviceAuth, oauthClient, logger, nil, ""), db: database, cfg: cfg, alice: caller{did: owner, signer: priv}, }, cfg.ServiceDid.String() } func TestCreateTaskWithADeadGrantRequiresReauthorization(t *testing.T) { g := newGrantStub(t, syntax.DID(alice)) g.dead = true ts, _ := serverWithGrants(t, g) path := "/xrpc/" + org_tangled.TempMigratorCreateTaskNSID body := &org_tangled.TempMigratorCreateTask_Input{RequestId: "req-dead", Jobs: []*org_tangled.TempMigratorDefs_NewJob{job("one", false)}} rec := ts.do(t, http.MethodPost, path, ts.token(t, ts.alice, org_tangled.TempMigratorCreateTaskNSID), body) if rec.Code != 428 { t.Fatalf("want 428 Precondition Required, got %d (body: %s)", rec.Code, rec.Body.String()) } if !strings.Contains(rec.Body.String(), "GrantRequired") { t.Fatalf("body = %s, want the GrantRequired tag", rec.Body.String()) } t.Logf("428 body: %s", rec.Body.String()) var batches, jobs int if err := ts.db.QueryRow("select count(*) from batches").Scan(&batches); err != nil { t.Fatal(err) } if err := ts.db.QueryRow("select count(*) from jobs").Scan(&jobs); err != nil { t.Fatal(err) } t.Logf("rows after dead-grant createTask: batches=%d jobs=%d", batches, jobs) if batches != 0 || jobs != 0 { t.Fatalf("a dead grant created %d batches and %d jobs", batches, jobs) } g.mu.Lock() before := g.requests g.mu.Unlock() repeat := ts.do(t, http.MethodPost, path, ts.token(t, ts.alice, org_tangled.TempMigratorCreateTaskNSID), body) if repeat.Code != 428 || !strings.Contains(repeat.Body.String(), "GrantRequired") { t.Fatalf("repeat createTask = %d %s, want the same 428 GrantRequired", repeat.Code, repeat.Body.String()) } g.mu.Lock() defer g.mu.Unlock() if g.requests != before { t.Fatalf("repeat createTask made %d stub requests, want 0 (the dead row must be gone)", g.requests-before) } } func TestCreateTaskWithAHealthyGrantStillCreatesTheTask(t *testing.T) { g := newGrantStub(t, syntax.DID(alice)) ts, serviceDid := serverWithGrants(t, g) path := "/xrpc/" + org_tangled.TempMigratorCreateTaskNSID body := &org_tangled.TempMigratorCreateTask_Input{RequestId: "req-alive", Jobs: []*org_tangled.TempMigratorDefs_NewJob{job("one", false), job("two", false)}} rec := ts.do(t, http.MethodPost, path, ts.token(t, ts.alice, org_tangled.TempMigratorCreateTaskNSID), body) if rec.Code != http.StatusAccepted { t.Fatalf("want 202, got %d (body: %s)", rec.Code, rec.Body.String()) } var batches, jobs int if err := ts.db.QueryRow("select count(*) from batches").Scan(&batches); err != nil { t.Fatal(err) } if err := ts.db.QueryRow("select count(*) from jobs").Scan(&jobs); err != nil { t.Fatal(err) } t.Logf("rows after healthy createTask: batches=%d jobs=%d", batches, jobs) if batches != 1 || jobs != 2 { t.Fatalf("a healthy grant produced batches=%d jobs=%d, want 1 and 2", batches, jobs) } g.mu.Lock() defer g.mu.Unlock() if g.refreshes != 1 { t.Fatalf("token refreshes = %d, want exactly the one the probe needed", g.refreshes) } if len(g.mintCalls) != 1 || g.mintCalls[0] != "sh.tangled.repo.describeRepo aud="+serviceDid { t.Fatalf("probe mint calls = %v, want one describeRepo mint for the migrator's own did", g.mintCalls) } }