From f78003ea11943d15bbec9d0d67968fd9afb7bff0 Mon Sep 17 00:00:00 2001 From: Amolith Date: Fri, 24 Jul 2026 20:51:33 -0600 Subject: [PATCH] app: guard record updates with compare-and-swap Record edits previously performed an unguarded read-modify-write, allowing a concurrent update to be overwritten. Carry the fetched CID into swapRecord so stale writes fail. Keep issue and pull edits on generated lexicon types, while repository edits preserve unknown fields and validate the modeled schema before writing. --- atproto/records.go | 15 ++--- atproto/records_test.go | 48 ++++++++++++++++ internal/app/records.go | 27 +++++++-- internal/app/records_test.go | 32 +++++++++++ internal/app/repos.go | 53 ++++++++---------- internal/app/repos_edit_test.go | 83 ++++++++++++++++++++++++++++ internal/tangledlex/validate.go | 14 +++++ internal/tangledlex/validate_test.go | 12 ++++ 8 files changed, 243 insertions(+), 41 deletions(-) create mode 100644 internal/app/repos_edit_test.go 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 -- 2.51.2