package oauth import ( "context" "crypto/rand" "encoding/base64" "encoding/json" "net/http" "net/http/httptest" "reflect" "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 "tangled.org/core/api/tangled" "tangled.org/core/log" "tangled.org/core/migrator/config" ) type serviceTokenCall struct { aud string lxm string } func TestCreateUsesUserServiceTokenAndDescribeUsesMigratorToken(t *testing.T) { store, database := testStore(t) userDID := syntax.DID("did:plc:alice") userKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } dpopKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } var mu sync.Mutex var pdsCalls []serviceTokenCall var knotClaims []map[string]any var server *httptest.Server server = httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { case "/oauth/token": if r.Header.Get("DPoP") == "" { t.Error("refresh omitted DPoP proof") } if err := r.ParseForm(); err != nil { t.Error(err) } if r.Form.Get("grant_type") != "refresh_token" { t.Errorf("refresh grant_type = %q", r.Form.Get("grant_type")) } _ = json.NewEncoder(w).Encode(indigooauth.TokenResponse{ Subject: userDID.String(), Scope: strings.Join(scopes, " "), AccessToken: "pds-oauth-refreshed", RefreshToken: "refresh-next", }) case "/xrpc/com.atproto.server.getServiceAuth": if r.Header.Get("Authorization") != "DPoP pds-oauth-old" || r.Header.Get("DPoP") == "" { t.Errorf("PDS auth headers = Authorization %q, DPoP present %t", r.Header.Get("Authorization"), r.Header.Get("DPoP") != "") } aud, lxm := r.URL.Query().Get("aud"), r.URL.Query().Get("lxm") parsed, err := syntax.ParseNSID(lxm) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } token, err := auth.SignServiceAuth(userDID, aud, time.Minute, &parsed, userKey) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } mu.Lock() pdsCalls = append(pdsCalls, serviceTokenCall{aud: aud, lxm: lxm}) mu.Unlock() _ = json.NewEncoder(w).Encode(map[string]string{"token": token}) case "/xrpc/sh.tangled.repo.create": claims, ok := bearerClaims(r.Header.Get("Authorization")) if !ok { t.Errorf("create knot authorization = %q", r.Header.Get("Authorization")) } mu.Lock() knotClaims = append(knotClaims, claims) mu.Unlock() repoDID := "did:plc:created" _ = json.NewEncoder(w).Encode(tangled.RepoCreate_Output{RepoDid: &repoDID}) case "/xrpc/sh.tangled.repo.describeRepo": claims, ok := bearerClaims(r.Header.Get("Authorization")) if !ok { t.Errorf("describe knot authorization = %q", r.Header.Get("Authorization")) } mu.Lock() knotClaims = append(knotClaims, claims) mu.Unlock() content := "present" _ = json.NewEncoder(w).Encode(tangled.RepoDescribeRepo_Output{ Content: &content, OwnerDid: userDID.String(), RepoDid: "did:plc:created", Rkey: "repo", }) default: http.NotFound(w, r) } })) defer server.Close() 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-knot-test")) if err != nil { t.Fatal(err) } client.app.Client = server.Client() client.http = server.Client() client.store = store client.app.Store = store if err := store.SaveSession(context.Background(), indigooauth.ClientSessionData{ AccountDID: userDID, SessionID: "session", HostURL: server.URL, AuthServerURL: server.URL, AuthServerTokenEndpoint: server.URL + "/oauth/token", Scopes: scopes, AccessToken: "pds-oauth-old", RefreshToken: "refresh-old", DPoPPrivateKeyMultibase: dpopKey.Multibase(), }); err != nil { t.Fatal(err) } host := strings.TrimPrefix(server.URL, "https://") knotDID := "did:web:" + strings.ReplaceAll(host, ":", "%3A") repoDID, err := client.CreateRepo(context.Background(), userDID.String(), knotDID, "repo", "repo", "https://migrator.example/git/cap/repo.git") if err != nil { t.Fatal(err) } if repoDID != "did:plc:created" { t.Fatalf("created repo = %q", repoDID) } content, _, err := client.DescribeRepo(context.Background(), userDID.String(), knotDID, repoDID) if err != nil || content != "present" { t.Fatalf("describe content = %q, err=%v", content, err) } polled, err := client.Content(context.Background(), knotDID, repoDID) if err != nil || polled != "present" { t.Fatalf("polled content = %q, err=%v", polled, err) } mu.Lock() defer mu.Unlock() wantCalls := []serviceTokenCall{ {aud: knotDID, lxm: tangled.RepoCreateNSID}, {aud: knotDID, lxm: tangled.RepoDescribeRepoNSID}, } if !reflect.DeepEqual(pdsCalls, wantCalls) { t.Fatalf("getServiceAuth calls = %+v, want %+v", pdsCalls, wantCalls) } if len(knotClaims) != 3 { t.Fatalf("knot claims = %+v", knotClaims) } if knotClaims[2]["iss"] != cfg.ServiceDid.String() || knotClaims[2]["aud"] != knotDID || knotClaims[2]["lxm"] != tangled.RepoDescribeRepoNSID { t.Fatalf("poll knot claims = %+v, want the migrator speaking for itself", knotClaims[2]) } if knotClaims[0]["iss"] != userDID.String() || knotClaims[0]["aud"] != knotDID || knotClaims[0]["lxm"] != tangled.RepoCreateNSID { t.Fatalf("create knot claims = %+v", knotClaims[0]) } if knotClaims[1]["iss"] != userDID.String() || knotClaims[1]["aud"] != knotDID || knotClaims[1]["lxm"] != tangled.RepoDescribeRepoNSID { t.Fatalf("describe knot claims = %+v", knotClaims[1]) } } func bearerClaims(header string) (map[string]any, bool) { if !strings.HasPrefix(header, "Bearer ") { return nil, false } parts := strings.Split(strings.TrimPrefix(header, "Bearer "), ".") if len(parts) != 3 { return nil, false } raw, err := base64.RawURLEncoding.DecodeString(parts[1]) if err != nil { return nil, false } var claims map[string]any if err := json.Unmarshal(raw, &claims); err != nil { return nil, false } return claims, true } func TestKnotHostIsTheDIDsHost(t *testing.T) { cases := []struct { did string want string ok bool }{ {"did:web:knot.example", "knot.example", true}, {"did:web:knot.tngl.boltless.dev", "knot.tngl.boltless.dev", true}, // a port is percent-escaped in a did:web host {"did:web:knot.example%3A8443", "knot.example:8443", true}, {"did:plc:abcdef", "", false}, {"knot.example", "", false}, } for _, tc := range cases { got, err := knotHost(tc.did) if tc.ok && (err != nil || got != tc.want) { t.Fatalf("knotHost(%q) = %q, %v; want %q", tc.did, got, err, tc.want) } if !tc.ok && err == nil { t.Fatalf("knotHost(%q) = %q, want an error", tc.did, got) } } } func TestRepoRecordMatchesTheAppShape(t *testing.T) { now := time.Date(2026, 9, 15, 12, 0, 0, 0, time.UTC) decode := func(record *tangled.Repo) map[string]any { t.Helper() raw, err := json.Marshal(record) if err != nil { t.Fatal(err) } var decoded map[string]any if err := json.Unmarshal(raw, &decoded); err != nil { t.Fatal(err) } return decoded } plain := decode(repoRecord("hello", "hello", " ", "knot.example", "did:plc:repo", now)) if plain["$type"] != tangled.RepoNSID || plain["createdAt"] != "2026-09-15T12:00:00Z" { t.Fatalf("record = %v", plain) } if plain["knot"] != "knot.example" || plain["repoDid"] != "did:plc:repo" { t.Fatalf("record = %v", plain) } if _, present := plain["name"]; present { t.Fatalf("record = %v, want no cosmetic name when it repeats the rkey", plain) } if _, present := plain["description"]; present { t.Fatalf("record = %v, want an empty description omitted", plain) } named := decode(repoRecord("myrepo", "MyRepo", "a tool", "knot.example", "did:plc:repo", now)) if named["name"] != "MyRepo" || named["description"] != "a tool" { t.Fatalf("record = %v, want the cosmetic name and description", named) } labels, ok := plain["labels"].([]any) if !ok || len(labels) != 5 { t.Fatalf("labels = %v, want the app's five default definitions", plain["labels"]) } if labels[0] != "at://did:plc:wshs7t2adsemcrrd4snkeqli/sh.tangled.label.definition/wontfix" { t.Fatalf("labels[0] = %v, want the app's first default label", labels[0]) } } type mockPdsRecord struct { cid string repoDid *string name string } func setupPutRecordTest(t *testing.T, initial *mockPdsRecord) (*Client, func() *mockPdsRecord, *[]string) { t.Helper() return setupPutRecordTestWithSwaps(t, initial, 0) } func setupPutRecordTestWithSwaps(t *testing.T, initial *mockPdsRecord, rejectSwaps int) (*Client, func() *mockPdsRecord, *[]string) { t.Helper() store, database := testStore(t) userDID := syntax.DID("did:plc:alice") userKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } dpopKey, err := atcrypto.GeneratePrivateKeyP256() if err != nil { t.Fatal(err) } var mu sync.Mutex currentRecord := initial var putCalls []string server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { case "/oauth/token": _ = json.NewEncoder(w).Encode(indigooauth.TokenResponse{ Subject: userDID.String(), Scope: strings.Join(scopes, " "), AccessToken: "pds-oauth-refreshed", RefreshToken: "refresh-next", }) case "/xrpc/com.atproto.server.getServiceAuth": aud, lxm := r.URL.Query().Get("aud"), r.URL.Query().Get("lxm") parsed, err := syntax.ParseNSID(lxm) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } token, err := auth.SignServiceAuth(userDID, aud, time.Minute, &parsed, userKey) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } _ = json.NewEncoder(w).Encode(map[string]string{"token": token}) case "/xrpc/com.atproto.repo.createRecord": mu.Lock() defer mu.Unlock() if currentRecord != nil { w.WriteHeader(http.StatusConflict) _ = json.NewEncoder(w).Encode(map[string]any{ "error": "InvalidRequest", "message": "record already exists", }) return } var body struct { Record map[string]any `json:"record"` } _ = json.NewDecoder(r.Body).Decode(&body) var did *string if d, ok := body.Record["repoDid"].(string); ok { did = &d } currentRecord = &mockPdsRecord{ cid: "bafy-created", repoDid: did, } w.WriteHeader(http.StatusOK) _ = json.NewEncoder(w).Encode(map[string]any{ "uri": "at://did:plc:alice/sh.tangled.repo/repo", "cid": "bafy-created", }) case "/xrpc/com.atproto.repo.getRecord": mu.Lock() defer mu.Unlock() if currentRecord == nil { w.WriteHeader(http.StatusNotFound) return } val := map[string]any{ "$type": tangled.RepoNSID, "knot": "knot.example", "name": currentRecord.name, "createdAt": "2026-09-15T12:00:00Z", } if currentRecord.repoDid != nil { val["repoDid"] = *currentRecord.repoDid } _ = json.NewEncoder(w).Encode(map[string]any{ "uri": "at://did:plc:alice/sh.tangled.repo/repo", "cid": currentRecord.cid, "value": val, }) case "/xrpc/com.atproto.repo.putRecord": mu.Lock() defer mu.Unlock() var body struct { Repo string `json:"repo"` Collection string `json:"collection"` Rkey string `json:"rkey"` SwapRecord *string `json:"swapRecord"` Record map[string]any `json:"record"` } if err := json.NewDecoder(r.Body).Decode(&body); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } putCalls = append(putCalls, body.Rkey) if rejectSwaps > 0 { rejectSwaps-- currentRecord = &mockPdsRecord{ cid: "bafy-moved", repoDid: currentRecord.repoDid, name: currentRecord.name, } w.WriteHeader(http.StatusBadRequest) _ = json.NewEncoder(w).Encode(map[string]any{ "error": "InvalidSwap", "message": "record has changed since it was read", }) return } var did *string if d, ok := body.Record["repoDid"].(string); ok { did = &d } currentRecord = &mockPdsRecord{ cid: "bafy-updated", repoDid: did, } w.WriteHeader(http.StatusOK) _ = json.NewEncoder(w).Encode(map[string]any{ "uri": "at://did:plc:alice/sh.tangled.repo/repo", "cid": "bafy-updated", }) default: http.NotFound(w, r) } })) t.Cleanup(server.Close) 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) } dir := identity.NewMockDirectory() dir.Insert(identity.Identity{ DID: userDID, Services: map[string]identity.ServiceEndpoint{ "atproto_pds": {Type: "AtprotoPersonalDataServer", URL: server.URL}, }, }) client, err := NewClient(cfg, database, dir, log.New("oauth-put-test")) if err != nil { t.Fatal(err) } client.app.Client = server.Client() client.http = server.Client() client.store = store client.app.Store = store if err := store.SaveSession(context.Background(), indigooauth.ClientSessionData{ AccountDID: userDID, SessionID: "session", HostURL: server.URL, AuthServerURL: server.URL, AuthServerTokenEndpoint: server.URL + "/oauth/token", Scopes: scopes, AccessToken: "pds-oauth-old", RefreshToken: "refresh-old", DPoPPrivateKeyMultibase: dpopKey.Multibase(), }); err != nil { t.Fatal(err) } getRecord := func() *mockPdsRecord { mu.Lock() defer mu.Unlock() return currentRecord } return client, getRecord, &putCalls } func TestPutRepoRecordOptimisticRecordWithoutRepoDidSucceeds(t *testing.T) { // existing record without repoDid succeeds and ends up naming the new repoDid client, getRecord, putCalls := setupPutRecordTest(t, &mockPdsRecord{cid: "bafy-initial", repoDid: nil, name: "repo"}) err := client.PutRepoRecord(context.Background(), "did:plc:alice", "repo", "repo", "desc", "did:web:knot.example", "did:plc:new-repo") if err != nil { t.Fatalf("PutRepoRecord error = %v, want success updating optimistic record", err) } current := getRecord() if current == nil || current.repoDid == nil || *current.repoDid != "did:plc:new-repo" { t.Fatalf("current record repoDid = %v, want did:plc:new-repo", current) } if len(*putCalls) != 1 { t.Fatalf("putRecord calls = %d, want 1", len(*putCalls)) } } func TestPutRepoRecordSameRepoDidIsIdempotent(t *testing.T) { // existing record naming the same repoDid succeeds without a second write sameDid := "did:plc:same" client, _, putCalls := setupPutRecordTest(t, &mockPdsRecord{cid: "bafy-initial", repoDid: &sameDid, name: "repo"}) err := client.PutRepoRecord(context.Background(), "did:plc:alice", "repo", "repo", "desc", "did:web:knot.example", "did:plc:same") if err != nil { t.Fatalf("PutRepoRecord error = %v, want idempotent success", err) } if len(*putCalls) != 0 { t.Fatalf("putRecord calls = %d, want 0 (no duplicate write)", len(*putCalls)) } } func TestPutRepoRecordDifferentRepoDidIsError(t *testing.T) { // existing record naming a different repoDid is still an error otherDid := "did:plc:other" client, _, _ := setupPutRecordTest(t, &mockPdsRecord{cid: "bafy-initial", repoDid: &otherDid, name: "repo"}) err := client.PutRepoRecord(context.Background(), "did:plc:alice", "repo", "repo", "desc", "did:web:knot.example", "did:plc:new") if err == nil || !strings.Contains(err.Error(), "already names repository did:plc:other") { t.Fatalf("PutRepoRecord error = %v, want error naming different repo", err) } } func TestPutRepoRecordRetriesALostSwap(t *testing.T) { client, getRecord, putCalls := setupPutRecordTestWithSwaps( t, &mockPdsRecord{cid: "bafy-initial", repoDid: nil, name: "repo"}, 1) err := client.PutRepoRecord(context.Background(), "did:plc:alice", "repo", "repo", "desc", "did:web:knot.example", "did:plc:new-repo") if err != nil { t.Fatalf("PutRepoRecord error = %v, want the retry to settle the record", err) } current := getRecord() if current == nil || current.repoDid == nil || *current.repoDid != "did:plc:new-repo" { t.Fatalf("record = %v, want it to name did:plc:new-repo after the retry", current) } if len(*putCalls) != 2 { t.Fatalf("putRecord attempts = %d, want 2 (one lost swap, one retry)", len(*putCalls)) } } func TestPutRepoRecordGivesUpAfterRepeatedLostSwaps(t *testing.T) { client, _, putCalls := setupPutRecordTestWithSwaps( t, &mockPdsRecord{cid: "bafy-initial", repoDid: nil, name: "repo"}, 9) err := client.PutRepoRecord(context.Background(), "did:plc:alice", "repo", "repo", "desc", "did:web:knot.example", "did:plc:new-repo") if err == nil { t.Fatal("PutRepoRecord error = nil, want a bounded failure rather than an endless retry") } if len(*putCalls) != 3 { t.Fatalf("putRecord attempts = %d, want the bound of 3", len(*putCalls)) } }