diff --git a/atproto/records.go b/atproto/records.go index 07ed402..9e4bfcc 100644 --- a/atproto/records.go +++ b/atproto/records.go @@ -21,10 +21,11 @@ type ATProto struct { } type PutRecordInput struct { - Repo string `json:"repo"` - Collection string `json:"collection"` - Rkey string `json:"rkey"` - Record any `json:"record"` + Repo string `json:"repo"` + Collection string `json:"collection"` + Rkey string `json:"rkey"` + Record any `json:"record"` + SwapRecord *syntax.CID `json:"swapRecord,omitempty"` } type DeleteRecordInput struct { @@ -42,9 +43,9 @@ type Blob struct { } type GetRecordOutput struct { - URI string `json:"uri"` - CID string `json:"cid,omitempty"` - Value any `json:"value"` + URI string `json:"uri"` + CID *syntax.CID `json:"cid,omitempty"` + Value any `json:"value"` } // RecordItem is a single record in a listRecords response. diff --git a/atproto/records_test.go b/atproto/records_test.go index 32cd75a..d4ac92c 100644 --- a/atproto/records_test.go +++ b/atproto/records_test.go @@ -61,3 +61,51 @@ func TestPutRecordSendsValidatedTangledRecord(t *testing.T) { t.Fatalf("PutRecord() URI = %q", uri) } } + +func TestRecordUpdateRoundTripsFetchedCIDAsSwapRecord(t *testing.T) { + var swapRecord string + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + writer.Header().Set("Content-Type", "application/json") + switch request.Method { + case http.MethodGet: + _, _ = writer.Write([]byte(`{"uri":"at://did:plc:abc123/sh.tangled.repo/example","cid":"bafyreifetched","value":{"$type":"sh.tangled.repo","knot":"knot.example","createdAt":"2026-07-25T12:00:00Z"}}`)) + case http.MethodPost: + defer request.Body.Close() + var input struct { + SwapRecord string `json:"swapRecord"` + } + if err := json.NewDecoder(request.Body).Decode(&input); err != nil { + t.Errorf("decode PutRecord request: %v", err) + return + } + swapRecord = input.SwapRecord + _, _ = writer.Write([]byte(`{"uri":"at://did:plc:abc123/sh.tangled.repo/example","cid":"bafyreinew"}`)) + default: + t.Errorf("request method = %s, want GET or POST", request.Method) + } + })) + defer server.Close() + + client := &ATProto{Client: &atclient.APIClient{Client: server.Client(), Host: server.URL}} + found, err := client.GetRecord(context.Background(), "did:plc:abc123", "sh.tangled.repo", "example") + if err != nil { + t.Fatalf("GetRecord() error = %v", err) + } + if found.CID == nil || found.CID.String() != "bafyreifetched" { + t.Fatalf("GetRecord() CID = %v, want bafyreifetched", found.CID) + } + + _, _, err = client.PutRecord(context.Background(), PutRecordInput{ + Repo: "did:plc:abc123", + Collection: "sh.tangled.repo", + Rkey: "example", + Record: found.Value, + SwapRecord: found.CID, + }) + if err != nil { + t.Fatalf("PutRecord() error = %v", err) + } + if swapRecord != "bafyreifetched" { + t.Fatalf("PutRecord() swapRecord = %q, want bafyreifetched", swapRecord) + } +} diff --git a/internal/app/records.go b/internal/app/records.go index 1fb2f86..3f27037 100644 --- a/internal/app/records.go +++ b/internal/app/records.go @@ -31,19 +31,38 @@ func putRecord(ctx context.Context, atClient pdsClient, did, collection, rkey st return nil } -// editRecord fetches an existing record, applies the provided title and/or -// body patches (nil leaves the field untouched), and writes it back. +// editRecord applies title and body patches to an existing record. func editRecord(ctx context.Context, atClient pdsClient, did, collection, rkey string, title, body *string) error { + return updateRecord(ctx, atClient, did, collection, rkey, func(value any) (any, error) { + return editLexiconRecord(collection, value, title, body) + }) +} + +// updateRecord fetches an existing record, applies mutate, and writes it back +// only if the fetched record is still current. +func updateRecord[T any](ctx context.Context, atClient pdsClient, did, collection, rkey string, mutate func(any) (T, error)) error { found, err := atClient.GetRecord(ctx, did, collection, rkey) if err != nil { return fmt.Errorf("get existing record: %w", err) } + if found.CID == nil { + return fmt.Errorf("get existing record: PDS response omitted record CID") + } - record, err := editLexiconRecord(collection, found.Value, title, body) + record, err := mutate(found.Value) if err != nil { return err } - return putRecord(ctx, atClient, did, collection, rkey, record) + if _, _, err := atClient.PutRecord(ctx, atproto.PutRecordInput{ + Repo: did, + Collection: collection, + Rkey: rkey, + Record: record, + SwapRecord: found.CID, + }); err != nil { + return err + } + return nil } func editLexiconRecord(collection string, value any, title, body *string) (any, error) { diff --git a/internal/app/records_test.go b/internal/app/records_test.go index a17011d..73d74cd 100644 --- a/internal/app/records_test.go +++ b/internal/app/records_test.go @@ -1,9 +1,11 @@ package app import ( + "context" "strings" "testing" + "github.com/alyraffauf/tg/atproto" "github.com/alyraffauf/tg/internal/tangledlex" ) @@ -46,3 +48,33 @@ func TestEditLexiconRecordRejectsHistoricalRecordAtWriteBoundary(t *testing.T) { t.Fatalf("ValidateRecord() error = %v, want type validation error", err) } } + +func TestEditRecordGuardsGeneratedUpdateWithFetchedCID(t *testing.T) { + pds := &testPDS{record: &atproto.GetRecordOutput{ + CID: mustParseCID(t, "bafyreifetched"), + Value: map[string]any{ + "$type": issueCollection, + "repo": "did:plc:abc123", + "title": "Old title", + "createdAt": "2026-07-25T12:00:00Z", + }, + }} + title := "New title" + + if err := editRecord(context.Background(), pds, "did:plc:owner", issueCollection, "issue", &title, nil); err != nil { + t.Fatalf("editRecord() error = %v", err) + } + if len(pds.puts) != 1 { + t.Fatalf("editRecord() writes = %d, want 1", len(pds.puts)) + } + if pds.puts[0].SwapRecord == nil || pds.puts[0].SwapRecord.String() != "bafyreifetched" { + t.Fatalf("editRecord() swapRecord = %v, want bafyreifetched", pds.puts[0].SwapRecord) + } + issue, ok := pds.puts[0].Record.(tangledlex.RepoIssue) + if !ok { + t.Fatalf("editRecord() record type = %T, want tangledlex.RepoIssue", pds.puts[0].Record) + } + if issue.Title != title { + t.Fatalf("editRecord() title = %q, want %q", issue.Title, title) + } +} diff --git a/internal/app/repos.go b/internal/app/repos.go index d6bfca6..b8c414b 100644 --- a/internal/app/repos.go +++ b/internal/app/repos.go @@ -289,38 +289,31 @@ func (s *Service) EditRepo(ctx context.Context, t Target, in EditRepoInput) (*Re return nil, err } rkey := extractRKey(repo.URI) - existing, err := atClient.GetRecord(ctx, did, repoCollection, rkey) - if err != nil { - return nil, fmt.Errorf("get repository record: %w", err) - } - record, err := repoRecordMap(existing.Value) - if err != nil { - return nil, err - } - if in.Description != nil { - record["description"] = *in.Description - } - if in.Website != nil { - record["website"] = *in.Website - } - if in.Spindle != nil { - record["spindle"] = *in.Spindle - } - if len(in.AddLabels) > 0 || len(in.RemoveLabels) > 0 { - labels := labelsFromRecord(record["labels"]) - for _, label := range in.AddLabels { - labels[label] = true + if err := updateRecord(ctx, atClient, did, repoCollection, rkey, func(value any) (map[string]any, error) { + record, err := repoRecordMap(value) + if err != nil { + return nil, err } - for _, label := range in.RemoveLabels { - delete(labels, label) + if in.Description != nil { + record["description"] = *in.Description } - record["labels"] = labelNames(labels) - } - if _, _, err := atClient.PutRecord(ctx, atproto.PutRecordInput{ - Repo: did, - Collection: repoCollection, - Rkey: rkey, - Record: record, + if in.Website != nil { + record["website"] = *in.Website + } + if in.Spindle != nil { + record["spindle"] = *in.Spindle + } + if len(in.AddLabels) > 0 || len(in.RemoveLabels) > 0 { + labels := labelsFromRecord(record["labels"]) + for _, label := range in.AddLabels { + labels[label] = true + } + for _, label := range in.RemoveLabels { + delete(labels, label) + } + record["labels"] = labelNames(labels) + } + return record, nil }); err != nil { return nil, fmt.Errorf("edit repository: %w", err) } diff --git a/internal/app/repos_edit_test.go b/internal/app/repos_edit_test.go new file mode 100644 index 0000000..1290cbc --- /dev/null +++ b/internal/app/repos_edit_test.go @@ -0,0 +1,83 @@ +package app + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/alyraffauf/tg/atproto" + "github.com/alyraffauf/tg/tangled" + "github.com/bluesky-social/indigo/atproto/syntax" +) + +func TestEditRepoGuardsRecordUpdateWithFetchedCID(t *testing.T) { + pds := &testPDS{record: &atproto.GetRecordOutput{ + CID: mustParseCID(t, "bafyreifetched"), + Value: map[string]any{"$type": repoCollection, "description": "old", "knot": "knot.example", "createdAt": "2026-07-25T12:00:00Z"}, + }} + service := editRepoTestService(pds) + description := "new" + + _, err := service.EditRepo(context.Background(), Target{Handle: "owner.test", Repo: "example"}, EditRepoInput{Description: &description}) + if err != nil { + t.Fatalf("EditRepo() error = %v", err) + } + if len(pds.puts) != 1 { + t.Fatalf("EditRepo() writes = %d, want 1", len(pds.puts)) + } + if pds.puts[0].SwapRecord == nil || pds.puts[0].SwapRecord.String() != "bafyreifetched" { + t.Fatalf("EditRepo() swapRecord = %v, want bafyreifetched", pds.puts[0].SwapRecord) + } +} + +func TestEditRepoRequiresFetchedCID(t *testing.T) { + pds := &testPDS{record: &atproto.GetRecordOutput{ + Value: map[string]any{"$type": repoCollection, "description": "old", "knot": "knot.example", "createdAt": "2026-07-25T12:00:00Z"}, + }} + service := editRepoTestService(pds) + description := "new" + + _, err := service.EditRepo(context.Background(), Target{Handle: "owner.test", Repo: "example"}, EditRepoInput{Description: &description}) + if err == nil || !strings.Contains(err.Error(), "omitted record CID") { + t.Fatalf("EditRepo() error = %v, want missing CID error", err) + } + if len(pds.puts) != 0 { + t.Fatalf("EditRepo() writes = %d, want 0", len(pds.puts)) + } +} + +func TestEditRepoSurfacesRecordUpdateConflict(t *testing.T) { + conflict := errors.New("InvalidSwap: record CID did not match") + pds := &testPDS{ + record: &atproto.GetRecordOutput{CID: mustParseCID(t, "bafyreistale"), Value: map[string]any{"$type": repoCollection, "knot": "knot.example", "createdAt": "2026-07-25T12:00:00Z"}}, + putErr: conflict, + } + service := editRepoTestService(pds) + description := "new" + + _, err := service.EditRepo(context.Background(), Target{Handle: "owner.test", Repo: "example"}, EditRepoInput{Description: &description}) + if err == nil || !errors.Is(err, conflict) || !strings.Contains(err.Error(), "edit repository") { + t.Fatalf("EditRepo() error = %v, want wrapped update conflict", err) + } + if len(pds.puts) != 1 { + t.Fatalf("EditRepo() writes = %d, want 1", len(pds.puts)) + } +} + +func editRepoTestService(pds *testPDS) *Service { + service := testService(pds, &testGit{}, &testKnot{}) + service.appview = testAppview{repo: &tangled.Repo{ + URI: "at://did:plc:owner/sh.tangled.repo/example", + }} + return service +} + +func mustParseCID(t *testing.T, raw string) *syntax.CID { + t.Helper() + cid, err := syntax.ParseCID(raw) + if err != nil { + t.Fatalf("parse test CID %q: %v", raw, err) + } + return &cid +} diff --git a/internal/tangledlex/validate.go b/internal/tangledlex/validate.go index 71b0987..4403342 100644 --- a/internal/tangledlex/validate.go +++ b/internal/tangledlex/validate.go @@ -1,6 +1,7 @@ package tangledlex import ( + "encoding/json" "fmt" "net/url" "time" @@ -32,6 +33,19 @@ func ValidateRecord(collection string, record any) error { return validatePublicKey(collection, value) case String: return validateString(collection, value) + case map[string]any: + if collection != "sh.tangled.repo" { + return fmt.Errorf("%s record has unsupported generated type %T", collection, record) + } + data, err := json.Marshal(value) + if err != nil { + return fmt.Errorf("encode %s record for validation: %w", collection, err) + } + var repo Repo + if err := json.Unmarshal(data, &repo); err != nil { + return fmt.Errorf("decode %s record for validation: %w", collection, err) + } + return validateRepo(collection, repo) default: return fmt.Errorf("%s record has unsupported generated type %T", collection, record) } diff --git a/internal/tangledlex/validate_test.go b/internal/tangledlex/validate_test.go index e656139..e9643f7 100644 --- a/internal/tangledlex/validate_test.go +++ b/internal/tangledlex/validate_test.go @@ -24,6 +24,18 @@ func TestValidateRecordAcceptsEveryWritableRecord(t *testing.T) { } } +func TestValidateRecordAcceptsPreservedRepoMap(t *testing.T) { + record := map[string]any{ + "$type": "sh.tangled.repo", + "knot": "knot.example", + "createdAt": testTime, + "custom": map[string]any{"preserved": true}, + } + if err := ValidateRecord("sh.tangled.repo", record); err != nil { + t.Fatalf("ValidateRecord() error = %v", err) + } +} + func TestValidateRecordRejectsEveryWritableRecord(t *testing.T) { tests := []struct { name string