diff --git a/atproto/client.go b/atproto/client.go index d8e31d6..bf83e0d 100644 --- a/atproto/client.go +++ b/atproto/client.go @@ -209,17 +209,6 @@ func (c *Client) APIClient() *atclient.APIClient { return c.client } -func (c *Client) Close() error { - c.mu.Lock() - defer c.mu.Unlock() - if c.client != nil && c.client.Auth != nil { - if logout, ok := c.client.Auth.(*atclient.PasswordAuth); ok { - return logout.Logout(context.Background(), c.client.Client) - } - } - return nil -} - func (c *Client) HasClient() bool { c.mu.Lock() defer c.mu.Unlock() diff --git a/atproto/client_test.go b/atproto/client_test.go index 16b33a0..842597f 100644 --- a/atproto/client_test.go +++ b/atproto/client_test.go @@ -6,6 +6,7 @@ import ( "net/http" "net/http/httptest" "testing" + "testing/synctest" "time" "github.com/bluesky-social/indigo/atproto/atclient" @@ -97,7 +98,7 @@ func TestResolveMiniDoc(t *testing.T) { })) defer server.Close() - ctx := context.Background() + ctx := t.Context() opts := tt.opts if opts == nil { @@ -144,7 +145,7 @@ func TestResolveMiniDoc_WithCustomResolver(t *testing.T) { opts := NewClientOptions() opts.ResolverURL = server.URL - ctx := context.Background() + ctx := t.Context() did, pds, _, err := ResolveMiniDoc(ctx, "test.user", opts) if err != nil { t.Fatalf("ResolveMiniDoc failed: %v", err) @@ -176,7 +177,7 @@ func TestResolveMiniDoc_WithUserAgent(t *testing.T) { opts.ResolverURL = server.URL opts.UserAgent = "test-agent/1.0" - ctx := context.Background() + ctx := t.Context() did, _, _, err := ResolveMiniDoc(ctx, "test.user", opts) if err != nil { t.Fatalf("ResolveMiniDoc failed: %v", err) @@ -188,32 +189,34 @@ func TestResolveMiniDoc_WithUserAgent(t *testing.T) { } func TestResolveMiniDoc_ContextCancelled(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - select { - case <-time.After(100 * time.Millisecond): - w.WriteHeader(http.StatusOK) - case <-r.Context().Done(): - return - } - })) - defer server.Close() + synctest.Test(t, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + select { + case <-time.After(100 * time.Millisecond): + w.WriteHeader(http.StatusOK) + case <-r.Context().Done(): + return + } + })) + defer server.Close() - ctx, cancel := context.WithCancel(context.Background()) - cancel() + ctx, cancel := context.WithCancel(t.Context()) + cancel() - opts := NewClientOptions() - opts.ResolverURL = server.URL - opts.HTTPClient = server.Client() + opts := NewClientOptions() + opts.ResolverURL = server.URL + opts.HTTPClient = server.Client() - _, _, _, err := ResolveMiniDoc(ctx, "test.user", opts) + _, _, _, err := ResolveMiniDoc(ctx, "test.user", opts) - if err == nil { - t.Error("expected error for cancelled context") - } + if err == nil { + t.Error("expected error for cancelled context") + } + }) } func TestResolveMiniDoc_InvalidURL(t *testing.T) { - ctx := context.Background() + ctx := t.Context() opts := NewClientOptions() opts.ResolverURL = "://invalid-url" @@ -281,7 +284,7 @@ func TestResolveIdentity(t *testing.T) { t.Parallel() if tt.handle == "" { - ctx := context.Background() + ctx := t.Context() _, err := ResolveIdentity(ctx, "", tt.opts) if err == nil { t.Error("expected error for empty handle") @@ -295,7 +298,7 @@ func TestResolveIdentity(t *testing.T) { })) defer server.Close() - ctx := context.Background() + ctx := t.Context() opts := tt.opts if opts == nil { @@ -325,132 +328,6 @@ func TestResolveIdentity(t *testing.T) { } } -func TestClient_Getters(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - setup func() *Client - check func(t *testing.T, c *Client) - }{ - { - name: "HasClient true", - setup: func() *Client { - return &Client{ - client: &atclient.APIClient{}, - resolvedIdentity: resolvedIdentity{ - DID: "did:plc:test", - Handle: "test.bsky.social", - PDS: "https://pds.example.com", - }, - } - }, - check: func(t *testing.T, c *Client) { - if !c.HasClient() { - t.Error("HasClient() = false, want true") - } - }, - }, - { - name: "HasClient false", - setup: func() *Client { - return &Client{} - }, - check: func(t *testing.T, c *Client) { - if c.HasClient() { - t.Error("HasClient() = true, want false") - } - }, - }, - { - name: "DID", - setup: func() *Client { - return &Client{ - resolvedIdentity: resolvedIdentity{DID: "did:plc:test123"}, - } - }, - check: func(t *testing.T, c *Client) { - if got := c.DID(); got != "did:plc:test123" { - t.Errorf("DID() = %s, want did:plc:test123", got) - } - }, - }, - { - name: "PDS", - setup: func() *Client { - return &Client{ - resolvedIdentity: resolvedIdentity{PDS: "https://pds.example.com"}, - } - }, - check: func(t *testing.T, c *Client) { - if got := c.PDS(); got != "https://pds.example.com" { - t.Errorf("PDS() = %s, want https://pds.example.com", got) - } - }, - }, - { - name: "Handle", - setup: func() *Client { - return &Client{ - resolvedIdentity: resolvedIdentity{Handle: "test.bsky.social"}, - } - }, - check: func(t *testing.T, c *Client) { - if got := c.Handle(); got != "test.bsky.social" { - t.Errorf("Handle() = %s, want test.bsky.social", got) - } - }, - }, - { - name: "SigningKey", - setup: func() *Client { - return &Client{ - resolvedIdentity: resolvedIdentity{SigningKey: "-----BEGIN PUBLIC KEY-----\nabc\n-----END PUBLIC KEY-----"}, - } - }, - check: func(t *testing.T, c *Client) { - if got := c.SigningKey(); got != "-----BEGIN PUBLIC KEY-----\nabc\n-----END PUBLIC KEY-----" { - t.Errorf("SigningKey() = %s, want expected key", got) - } - }, - }, - { - name: "APIClient", - setup: func() *Client { - expectedClient := &atclient.APIClient{} - return &Client{ - client: expectedClient, - resolvedIdentity: resolvedIdentity{DID: "did:plc:test"}, - } - }, - check: func(t *testing.T, c *Client) { - if got := c.APIClient(); got == nil { - t.Error("APIClient() = nil, want non-nil") - } - }, - }, - { - name: "APIClient nil", - setup: func() *Client { - return &Client{} - }, - check: func(t *testing.T, c *Client) { - if got := c.APIClient(); got != nil { - t.Errorf("APIClient() = %v, want nil", got) - } - }, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - c := tt.setup() - tt.check(t, c) - }) - } -} - func TestBuildClient(t *testing.T) { t.Parallel() diff --git a/atproto/rate_test.go b/atproto/rate_test.go index 806f69d..d3ef0b9 100644 --- a/atproto/rate_test.go +++ b/atproto/rate_test.go @@ -53,7 +53,7 @@ func TestQuotaLimiter_AllowRead(t *testing.T) { rlQuota: 1.0, } - ctx := context.Background() + ctx := t.Context() chargedAt, err := rl.AllowRead(ctx) if err != nil { t.Fatalf("AllowRead failed: %v", err) @@ -98,7 +98,7 @@ func TestQuotaLimiter_AllowBulkWrite(t *testing.T) { rlQuota: 1.0, } - ctx := context.Background() + ctx := t.Context() _, err := rl.AllowBulkWrite(ctx, tt.n) if tt.wantErr && err == nil { @@ -167,7 +167,7 @@ func TestQuotaLimiter_Refund(t *testing.T) { rlQuota: 1.0, } - ctx := context.Background() + ctx := t.Context() tt.refundFunc(rl, ctx) if len(kv.incrs) != tt.wantIncrCalls { @@ -548,7 +548,7 @@ func TestQuotaLimiter_KVErrors(t *testing.T) { { name: "AllowBulkWrite error", testFunc: func(t *testing.T, rl *quotaLimiter) { - ctx := context.Background() + ctx := t.Context() _, err := rl.AllowBulkWrite(ctx, 1) if err == nil { t.Error("expected error from KV store") @@ -558,7 +558,7 @@ func TestQuotaLimiter_KVErrors(t *testing.T) { { name: "AllowRead error", testFunc: func(t *testing.T, rl *quotaLimiter) { - ctx := context.Background() + ctx := t.Context() _, err := rl.AllowRead(ctx) if err == nil { t.Error("expected error from KV store") @@ -651,7 +651,7 @@ func TestQuotaLimiter_Refund_KVErrors(t *testing.T) { rlQuota: 1.0, } - ctx := context.Background() + ctx := t.Context() tt.refundFn(rl, ctx) }) } @@ -668,7 +668,7 @@ func TestQuotaLimiter_MutexProtection(t *testing.T) { rlQuota: 1.0, } - ctx := context.Background() + ctx := t.Context() done := make(chan bool) errors := make(chan error, 10) diff --git a/atproto/repo_test.go b/atproto/repo_test.go index 49d3e11..5dd3117 100644 --- a/atproto/repo_test.go +++ b/atproto/repo_test.go @@ -169,7 +169,7 @@ func TestRateClient_ListRecords(t *testing.T) { } rateClient := NewRateClient[map[string]any](tt.client, "did:plc:test", tt.limiter) - ctx := context.Background() + ctx := t.Context() records, cursor, err := rateClient.ListRecords(ctx, tt.collection, tt.limit, tt.cursor) @@ -207,7 +207,7 @@ func TestRateClient_ListRecords_NetworkTimeout(t *testing.T) { } rateClient := NewRateClient[map[string]any](&client, "did:plc:test", nil) - ctx := context.Background() + ctx := t.Context() _, _, err := rateClient.ListRecords(ctx, "app.bsky.feed.post", 10, "") @@ -233,7 +233,7 @@ func TestRateClient_ListRecords_ContextCancelled(t *testing.T) { } rateClient := NewRateClient[map[string]any](&client, "did:plc:test", nil) - ctx, cancel := context.WithCancel(context.Background()) + ctx, cancel := context.WithCancel(t.Context()) cancel() _, _, err := rateClient.ListRecords(ctx, "app.bsky.feed.post", 10, "") @@ -277,7 +277,7 @@ func TestRateClient_ListRecords_WithPagination(t *testing.T) { } rateClient := NewRateClient[map[string]any](&client, "did:plc:test", nil) - ctx := context.Background() + ctx := t.Context() records, cursor, err := rateClient.ListRecords(ctx, "app.bsky.feed.post", 10, "") if err != nil { @@ -394,7 +394,7 @@ func TestRateClient_ApplyWrites(t *testing.T) { } rateClient := NewRateClient[map[string]any](tt.client, "did:plc:test", nil) - ctx := context.Background() + ctx := t.Context() err := rateClient.ApplyWrites(ctx, "app.bsky.feed.post", tt.records) @@ -421,7 +421,7 @@ func TestRateClient_ApplyWrites_NetworkTimeout(t *testing.T) { } rateClient := NewRateClient[map[string]any](&client, "did:plc:test", nil) - ctx := context.Background() + ctx := t.Context() records := []map[string]any{ {"$type": "app.bsky.feed.post", "text": "Hello"}, @@ -449,7 +449,7 @@ func TestRateClient_ApplyWrites_WithLimiter(t *testing.T) { limiter := NewRateLimiter(mockKV, 1.0) rateClient := NewRateClient[map[string]any](&client, "did:plc:test", limiter) - ctx := context.Background() + ctx := t.Context() records := []map[string]any{ {"$type": "app.bsky.feed.post", "text": "Hello"}, @@ -544,7 +544,7 @@ func TestRateClient_DeleteRecord(t *testing.T) { } rateClient := NewRateClient[map[string]any](tt.client, "did:plc:test", nil) - ctx := context.Background() + ctx := t.Context() err := rateClient.DeleteRecord(ctx, "app.bsky.feed.post", "3k5x3x2x1") @@ -571,7 +571,7 @@ func TestRateClient_DeleteRecord_NetworkTimeout(t *testing.T) { } rateClient := NewRateClient[map[string]any](&client, "did:plc:test", nil) - ctx := context.Background() + ctx := t.Context() err := rateClient.DeleteRecord(ctx, "app.bsky.feed.post", "3k5x3x2x1") @@ -634,7 +634,7 @@ func TestRepoClientFuncs(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - ctx := context.Background() + ctx := t.Context() switch tt.name { case "ListRecords": @@ -752,7 +752,7 @@ func TestPrepareWrites(t *testing.T) { } func TestApplyWrites_Empty(t *testing.T) { - err := applyWrites(context.Background(), nil, "did:plc:test", "app.bsky.feed.post", []map[string]any{}) + err := applyWrites(t.Context(), nil, "did:plc:test", "app.bsky.feed.post", []map[string]any{}) if err != nil { t.Errorf("applyWrites with empty records should not fail: %v", err) } @@ -937,7 +937,7 @@ func TestRateClient_ListRecords_WithTypedRecords(t *testing.T) { } rateClient := NewRateClient[testRecord](&client, "did:plc:test", nil) - ctx := context.Background() + ctx := t.Context() records, _, err := rateClient.ListRecords(ctx, "app.bsky.feed.post", 10, "") if err != nil { diff --git a/cache/bbolt.go b/cache/bbolt.go index 619aeb7..c3467a4 100644 --- a/cache/bbolt.go +++ b/cache/bbolt.go @@ -1,9 +1,11 @@ package cache import ( + "bytes" "encoding/json" "errors" "fmt" + "iter" "os" "path/filepath" "strings" @@ -22,7 +24,7 @@ type BoltStorage struct { path string } -func NewBoltStorage() (*BoltStorage, error) { +func NewBoltStorage(readOnly bool) (*BoltStorage, error) { dir, err := cacheDir() if err != nil { return nil, err @@ -31,7 +33,8 @@ func NewBoltStorage() (*BoltStorage, error) { return nil, err } db, err := bbolt.Open(filepath.Join(dir, CacheFile), 0o644, &bbolt.Options{ - Timeout: time.Second, + Timeout: time.Second, + ReadOnly: readOnly, }) if err != nil { return nil, err @@ -66,68 +69,117 @@ func (s *BoltStorage) SaveRecords(did string, records map[string][]byte) error { }) } -func (s *BoltStorage) IterateUnpublished(did string, fn func(key string, rec []byte) error) error { - published, err := s.GetPublished(did) - if err != nil { - return err - } - - return s.db.View(func(tx *bbolt.Tx) error { - b := tx.Bucket([]byte(recordsBucket(did))) - if b == nil { - return nil +func (s *BoltStorage) IterateUnpublished(did string, reverse bool) iter.Seq2[string, []byte] { + return func(yield func(key string, rec []byte) bool) { + published, err := s.GetPublished(did) + if err != nil { + return } - return b.ForEach(func(k, v []byte) error { - key := string(k) - if published[key] { + _ = s.db.View(func(tx *bbolt.Tx) error { + b := tx.Bucket([]byte(recordsBucket(did))) + if b == nil { return nil } - return fn(key, v) - }) - }) -} -func (s *BoltStorage) IteratePublished(did string, fn func(key string, rec []byte) error) error { - published, err := s.GetPublished(did) - if err != nil { - return err + if reverse { + cursor := b.Cursor() + + for k, v := cursor.Last(); k != nil; k, v = cursor.Prev() { + key := string(k) + if !published[key] { + // Clone key and value to avoid holding read locks + if !yield(key, bytes.Clone(v)) { + break // Stop iteration when yield returns false + } + } + } + } else { + return b.ForEach(func(k, v []byte) error { + key := string(k) + if published[key] { + return nil + } + // Clone key and value to avoid holding read locks + if !yield(key, bytes.Clone(v)) { + return errors.New("stop iteration") + } + return nil + }) + } + return nil + }) } +} - return s.db.View(func(tx *bbolt.Tx) error { - b := tx.Bucket([]byte(recordsBucket(did))) - if b == nil { - return nil +func (s *BoltStorage) IteratePublished(did string, reverse bool) iter.Seq2[string, []byte] { + return func(yield func(key string, rec []byte) bool) { + published, err := s.GetPublished(did) + if err != nil { + return } - return b.ForEach(func(k, v []byte) error { - key := string(k) - if !published[key] { + _ = s.db.View(func(tx *bbolt.Tx) error { + b := tx.Bucket([]byte(recordsBucket(did))) + if b == nil { return nil } - return fn(key, v) + + if reverse { + cursor := b.Cursor() + k, v := cursor.Last() + for k != nil { + key := string(k) + if published[key] { + // Clone key and value to avoid holding read locks + if !yield(key, bytes.Clone(v)) { + break // Stop iteration when yield returns false + } + } + k, v = cursor.Prev() + } + } else { + return b.ForEach(func(k, v []byte) error { + key := string(k) + if !published[key] { + return nil + } + // Clone key and value to avoid holding read locks + if !yield(key, bytes.Clone(v)) { + return errors.New("stop iteration") + } + return nil + }) + } + return nil }) - }) + } } -func (s *BoltStorage) IterateFailed(did string, fn func(key string, rec []byte, errMsg string) error) error { - return s.db.View(func(tx *bbolt.Tx) error { - fb := tx.Bucket([]byte(failedBucket(did))) - if fb == nil { - return nil - } - rb := tx.Bucket([]byte(recordsBucket(did))) - if rb == nil { - return nil - } +func (s *BoltStorage) IterateFailed(did string) func(yield func(key string, rec []byte, errMsg string) bool) { + return func(yield func(key string, rec []byte, errMsg string) bool) { + _ = s.db.View(func(tx *bbolt.Tx) error { + fb := tx.Bucket([]byte(failedBucket(did))) + if fb == nil { + return nil + } + rb := tx.Bucket([]byte(recordsBucket(did))) + if rb == nil { + return nil + } - return fb.ForEach(func(k, v []byte) error { - key := string(k) - errMsg := string(v) - rec := rb.Get(k) - return fn(key, rec, errMsg) + return fb.ForEach(func(k, v []byte) error { + key := string(k) + errMsg := string(v) + rec := rb.Get(k) + // Clone to avoid holding read lock + if !yield(key, bytes.Clone(rec), errMsg) { + return errors.New("stop iteration") + } + return nil + }) }) - }) + } } func (s *BoltStorage) MarkPublished(did string, keys ...string) error { @@ -302,45 +354,6 @@ func (s *BoltStorage) ClearAll() error { }) } -func (s *BoltStorage) Get(key string) (int, error) { - var val int - err := s.db.View(func(tx *bbolt.Tx) error { - b := tx.Bucket([]byte("quota")) - if b == nil { - return nil - } - v := b.Get([]byte(key)) - if v == nil { - return nil - } - return json.Unmarshal(v, &val) - }) - return val, err -} - -func (s *BoltStorage) IncrBy(key string, n int) (int, error) { - var val int - err := s.db.Update(func(tx *bbolt.Tx) error { - b, err := tx.CreateBucketIfNotExists([]byte("quota")) - if err != nil { - return err - } - v := b.Get([]byte(key)) - if v != nil { - if err := json.Unmarshal(v, &val); err != nil { - return err - } - } - val += n - newV, err := json.Marshal(val) - if err != nil { - return err - } - return b.Put([]byte(key), newV) - }) - return val, err -} - func (s *BoltStorage) GetMulti(keys []string) (map[string]int, error) { res := make(map[string]int, len(keys)) err := s.db.View(func(tx *bbolt.Tx) error { diff --git a/cache/cache_test.go b/cache/cache_test.go index b85c143..7e9ca3d 100644 --- a/cache/cache_test.go +++ b/cache/cache_test.go @@ -1,22 +1,93 @@ package cache import ( + "slices" "testing" ) func newTestStorage(t *testing.T) *BoltStorage { t.Helper() - storage, err := NewBoltStorage() + storage, err := NewBoltStorage(false) if err != nil { t.Fatalf("NewBoltStorage failed: %v", err) } - t.Cleanup(func() { storage.Close() }) + t.Cleanup(func() { + storage.Close() + }) return storage } +// testDID generates a unique DID for each test to ensure test isolation. +func testDID(t *testing.T, suffix string) string { + return "did:plc:test/" + t.Name() + "/" + suffix +} + +func TestReadOnlyMode(t *testing.T) { + tests := []struct { + name string + readOnly bool + shouldWrite bool + expectError bool + }{ + { + name: "read-write mode allows writes", + readOnly: false, + shouldWrite: true, + expectError: false, + }, + { + name: "read-only mode prevents writes", + readOnly: true, + shouldWrite: true, + expectError: true, + }, + { + name: "read-only mode allows iteration", + readOnly: true, + shouldWrite: false, + expectError: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + storage, err := NewBoltStorage(tt.readOnly) + if err != nil { + t.Fatalf("NewBoltStorage failed: %v", err) + } + t.Cleanup(func() { + storage.Close() + }) + + did := "did:plc:testro" + t.Name() + records := map[string][]byte{ + "key1": []byte(`{"trackName":"track1"}`), + } + + if tt.shouldWrite { + err = storage.SaveRecords(did, records) + if tt.expectError && err == nil { + t.Error("expected error when writing to read-only storage") + } else if !tt.expectError && err != nil { + t.Errorf("unexpected error when writing: %v", err) + } + } else { + // Just test that iteration works on read-only storage + var count int + for range storage.IterateUnpublished(did, false) { + count++ + } + if count != 0 { + t.Errorf("expected 0 records, got %d", count) + } + } + }) + } +} + func TestSaveIterateRoundtrip(t *testing.T) { storage := newTestStorage(t) - did := "did:plc:test" + did := testDID(t, "roundtrip") records := map[string][]byte{ "key1": []byte(`{"trackName":"track1"}`), @@ -28,12 +99,8 @@ func TestSaveIterateRoundtrip(t *testing.T) { } count := 0 - err := storage.IterateUnpublished(did, func(key string, data []byte) error { + for range storage.IterateUnpublished(did, false) { count++ - return nil - }) - if err != nil { - t.Fatalf("IterateUnpublished failed: %v", err) } if count != 2 { t.Errorf("expected 2 records, got %d", count) @@ -42,7 +109,7 @@ func TestSaveIterateRoundtrip(t *testing.T) { func TestMarkPublished(t *testing.T) { storage := newTestStorage(t) - did := "did:plc:test" + did := testDID(t, "mark") records := map[string][]byte{ "key1": []byte(`{"trackName":"track1"}`), @@ -55,13 +122,9 @@ func TestMarkPublished(t *testing.T) { } count := 0 - storage.IterateUnpublished(did, func(key string, data []byte) error { - if key == "key1" { - t.Error("key1 should have been filtered out") - } + for range storage.IterateUnpublished(did, false) { count++ - return nil - }) + } if count != 1 { t.Errorf("expected 1 unpublished record, got %d", count) } @@ -69,7 +132,7 @@ func TestMarkPublished(t *testing.T) { func TestClear(t *testing.T) { storage := newTestStorage(t) - did := "did:plc:test" + did := testDID(t, "clear") storage.SaveRecords(did, map[string][]byte{"key1": []byte(`{}`)}) if !storage.IsValid(did) { @@ -82,3 +145,346 @@ func TestClear(t *testing.T) { t.Error("cache should be invalid") } } + +func TestIterateUnpublishedReverse(t *testing.T) { + tests := []struct { + name string + did string + records map[string][]byte + reverse bool + expectedKeys []string + }{ + { + name: "basic reverse iteration", + records: map[string][]byte{"aaa": []byte(`{"trackName":"a"}`), "bbb": []byte(`{"trackName":"b"}`), "ccc": []byte(`{"trackName":"c"}`)}, + reverse: true, + expectedKeys: []string{"ccc", "bbb", "aaa"}, + }, + { + name: "forward iteration", + records: map[string][]byte{"aaa": []byte(`{"trackName":"a"}`), "bbb": []byte(`{"trackName":"b"}`), "ccc": []byte(`{"trackName":"c"}`)}, + reverse: false, + expectedKeys: []string{"aaa", "bbb", "ccc"}, + }, + { + name: "single record reverse", + records: map[string][]byte{"only": []byte(`{"trackName":"one"}`)}, + reverse: true, + expectedKeys: []string{"only"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + storage := newTestStorage(t) + did := testDID(t, "unpublished") + + if err := storage.SaveRecords(did, tt.records); err != nil { + t.Fatalf("SaveRecords failed: %v", err) + } + + var keys []string + for key, rec := range storage.IterateUnpublished(did, tt.reverse) { + keys = append(keys, key) + _ = rec + } + + if !slices.Equal(keys, tt.expectedKeys) { + t.Errorf("expected keys %v, got %v", tt.expectedKeys, keys) + } + }) + } +} + +func TestIteratePublishedReverse(t *testing.T) { + tests := []struct { + name string + records map[string][]byte + published []string + reverse bool + expectedKeys []string + }{ + { + name: "basic published reverse", + records: map[string][]byte{"aaa": []byte(`{"trackName":"a"}`), "bbb": []byte(`{"trackName":"b"}`), "ccc": []byte(`{"trackName":"c"}`)}, + published: []string{"aaa", "ccc"}, reverse: true, expectedKeys: []string{"ccc", "aaa"}, + }, + { + name: "forward published", + records: map[string][]byte{"aaa": []byte(`{"trackName":"a"}`), "bbb": []byte(`{"trackName":"b"}`), "ccc": []byte(`{"trackName":"c"}`)}, + published: []string{"aaa", "ccc"}, reverse: false, expectedKeys: []string{"aaa", "ccc"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + storage := newTestStorage(t) + did := testDID(t, "published") + + if err := storage.SaveRecords(did, tt.records); err != nil { + t.Fatalf("SaveRecords failed: %v", err) + } + + if err := storage.MarkPublished(did, tt.published...); err != nil { + t.Fatalf("MarkPublished failed: %v", err) + } + + var keys []string + for key, rec := range storage.IteratePublished(did, tt.reverse) { + keys = append(keys, key) + _ = rec + } + + if !slices.Equal(keys, tt.expectedKeys) { + t.Errorf("expected keys %v, got %v", tt.expectedKeys, keys) + } + }) + } +} + +func TestIterateReverseEmptyBucket(t *testing.T) { + tests := []struct { + name string + reverse bool + expectedLen int + }{ + { + name: "empty unpublished reverse", + reverse: true, + expectedLen: 0, + }, + { + name: "empty published reverse", + reverse: true, + expectedLen: 0, + }, + { + name: "empty unpublished forward", + reverse: false, + expectedLen: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + storage := newTestStorage(t) + did := testDID(t, "empty") + + count := 0 + for range storage.IterateUnpublished(did, tt.reverse) { + count++ + } + + if count != tt.expectedLen { + t.Errorf("expected %d records, got %d", tt.expectedLen, count) + } + }) + } +} + +func TestIterateReverseWithEarlyExit(t *testing.T) { + tests := []struct { + name string + records map[string][]byte + reverse bool + breakAfter int + expectedKeys []string + }{ + { + name: "exit after first record reverse", + records: map[string][]byte{"aaa": []byte(`{"trackName":"a"}`), "bbb": []byte(`{"trackName":"b"}`), "ccc": []byte(`{"trackName":"c"}`)}, + reverse: true, + breakAfter: 1, + expectedKeys: []string{"ccc"}, + }, + { + name: "exit after two records reverse", + records: map[string][]byte{"aaa": []byte(`{"trackName":"a"}`), "bbb": []byte(`{"trackName":"b"}`), "ccc": []byte(`{"trackName":"c"}`)}, + reverse: true, + breakAfter: 2, + expectedKeys: []string{"ccc", "bbb"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + storage := newTestStorage(t) + did := testDID(t, "exit") + + if err := storage.SaveRecords(did, tt.records); err != nil { + t.Fatalf("SaveRecords failed: %v", err) + } + + var keys []string + for key, rec := range storage.IterateUnpublished(did, tt.reverse) { + keys = append(keys, key) + _ = rec + if len(keys) >= tt.breakAfter { + break + } + } + + if !slices.Equal(keys, tt.expectedKeys) { + t.Errorf("expected keys %v, got %v", tt.expectedKeys, keys) + } + }) + } +} + +func TestIterateFailed(t *testing.T) { + tests := []struct { + name string + setupStorage func(*BoltStorage, string) error + wantCount int + wantKeys []string + }{ + { + name: "returns failed records", + setupStorage: func(s *BoltStorage, did string) error { + records := map[string][]byte{ + "key1": []byte(`{"trackName":"a"}`), + "key2": []byte(`{"trackName":"b"}`), + "key3": []byte(`{"trackName":"c"}`), + } + if err := s.SaveRecords(did, records); err != nil { + return err + } + return s.MarkFailed(did, []string{"key1", "key3"}, "timeout error") + }, + wantCount: 2, + wantKeys: []string{"key1", "key3"}, + }, + { + name: "handles empty failed set", + setupStorage: func(s *BoltStorage, did string) error { + records := map[string][]byte{ + "key1": []byte(`{"trackName":"a"}`), + } + return s.SaveRecords(did, records) + }, + wantCount: 0, + wantKeys: nil, + }, + { + name: "handles non-existent DID", + setupStorage: func(s *BoltStorage, did string) error { + return nil + }, + wantCount: 0, + wantKeys: nil, + }, + { + name: "returns error messages", + setupStorage: func(s *BoltStorage, did string) error { + records := map[string][]byte{ + "key1": []byte(`{"trackName":"a"}`), + } + if err := s.SaveRecords(did, records); err != nil { + return err + } + return s.MarkFailed(did, []string{"key1"}, "custom error message") + }, + wantCount: 1, + wantKeys: []string{"key1"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + storage := newTestStorage(t) + did := testDID(t, "failed") + + if err := tt.setupStorage(storage, did); err != nil { + t.Fatalf("setupStorage failed: %v", err) + } + + var count int + var keys []string + iterateFailed := storage.IterateFailed(did) + iterateFailed(func(key string, rec []byte, errMsg string) bool { + count++ + keys = append(keys, key) + return true + }) + + if count != tt.wantCount { + t.Errorf("IterateFailed() returned %d records, want %d", count, tt.wantCount) + } + + if len(keys) != len(tt.wantKeys) { + t.Errorf("IterateFailed() returned %d keys, want %d", len(keys), len(tt.wantKeys)) + return + } + + for i, key := range keys { + if key != tt.wantKeys[i] { + t.Errorf("IterateFailed() key[%d] = %s, want %s", i, key, tt.wantKeys[i]) + } + } + }) + } +} + +func TestIterateFailedWithEarlyExit(t *testing.T) { + storage := newTestStorage(t) + did := testDID(t, "failed-exit") + + records := map[string][]byte{ + "key1": []byte(`{"trackName":"a"}`), + "key2": []byte(`{"trackName":"b"}`), + "key3": []byte(`{"trackName":"c"}`), + } + if err := storage.SaveRecords(did, records); err != nil { + t.Fatalf("SaveRecords failed: %v", err) + } + if err := storage.MarkFailed(did, []string{"key1", "key2", "key3"}, "error"); err != nil { + t.Fatalf("MarkFailed failed: %v", err) + } + + // Exit after 2 records + var count int + iterateFailed := storage.IterateFailed(did) + iterateFailed(func(key string, rec []byte, errMsg string) bool { + count++ + return count < 2 // Exit after 2 records + }) + + if count != 2 { + t.Errorf("IterateFailed() returned %d records with early exit, want 2", count) + } +} + +func TestIterateFailedPreservesRecordData(t *testing.T) { + storage := newTestStorage(t) + did := testDID(t, "failed-data") + + records := map[string][]byte{ + "key1": []byte(`{"trackName":"Test Track","artist":"Test Artist"}`), + } + if err := storage.SaveRecords(did, records); err != nil { + t.Fatalf("SaveRecords failed: %v", err) + } + if err := storage.MarkFailed(did, []string{"key1"}, "network error"); err != nil { + t.Fatalf("MarkFailed failed: %v", err) + } + + var gotRecord []byte + var gotErrMsg string + iterateFailed := storage.IterateFailed(did) + iterateFailed(func(key string, rec []byte, errMsg string) bool { + if key == "key1" { + gotRecord = rec + gotErrMsg = errMsg + } + return true + }) + + expectedRecord := []byte(`{"trackName":"Test Track","artist":"Test Artist"}`) + if string(gotRecord) != string(expectedRecord) { + t.Errorf("IterateFailed() record = %s, want %s", string(gotRecord), string(expectedRecord)) + } + + if gotErrMsg != "network error" { + t.Errorf("IterateFailed() errMsg = %s, want 'network error'", gotErrMsg) + } +} diff --git a/cache/storage.go b/cache/storage.go index 50c2a0a..bebc5d3 100644 --- a/cache/storage.go +++ b/cache/storage.go @@ -1,34 +1,59 @@ package cache import ( + "iter" "time" ) -type Storage interface { +// RecordStore handles record persistence operations. +// Iterator methods use iter.Seq2 for iteration, with keys and values +// cloned to memory to avoid holding read locks during iteration. +type RecordStore interface { + // SaveRecords stores records for a given DID. SaveRecords(did string, records map[string][]byte) error - IterateUnpublished(did string, fn func(key string, rec []byte) error) error - IteratePublished(did string, fn func(key string, rec []byte) error) error - IterateFailed(did string, fn func(key string, rec []byte, errMsg string) error) error + + // IterateUnpublished iterates over unpublished records. + // The iterator yields (key, record) pairs. + // If reverse is true, iterates in reverse order using Cursor.Last/Cursor.Prev. + IterateUnpublished(did string, reverse bool) iter.Seq2[string, []byte] + + // IteratePublished iterates over published records. + // The iterator yields (key, record) pairs. + // If reverse is true, iterates in reverse order using Cursor.Last/Cursor.Prev. + IteratePublished(did string, reverse bool) iter.Seq2[string, []byte] + + // IterateFailed iterates over failed records. + // Returns a function that yields (key, record, errorMessage) triples. + IterateFailed(did string) func(yield func(key string, rec []byte, errMsg string) bool) + MarkPublished(did string, keys ...string) error MarkFailed(did string, keys []string, err string) error RemoveFailed(did string, keys ...string) error GetPublished(did string) (map[string]bool, error) + Clear(did string) error + Close() error +} + +// QuotaStore handles rate limit quota tracking. +type QuotaStore interface { + GetMulti(keys []string) (map[string]int, error) + IncrByMulti(counts map[string]int) error +} + +// Storage combines RecordStore and QuotaStore interfaces. +// Provided for backwards compatibility. +type Storage interface { + RecordStore + QuotaStore + IsValid(did string) bool Timestamp(did string) (time.Time, error) - Clear(did string) error ClearAll() error - Close() error // Stats returns database statistics Stats() (DBStats, error) - - // KVStore implementation - Get(key string) (int, error) - IncrBy(key string, n int) (int, error) - GetMulti(keys []string) (map[string]int, error) - IncrByMulti(counts map[string]int) error } type DBStats struct { diff --git a/flake.lock b/flake.lock index 8371674..aa33e0c 100644 --- a/flake.lock +++ b/flake.lock @@ -2,14 +2,16 @@ "nodes": { "flake-parts": { "inputs": { - "nixpkgs-lib": "nixpkgs-lib" + "nixpkgs-lib": [ + "nixpkgs" + ] }, "locked": { - "lastModified": 1768135262, - "narHash": "sha256-PVvu7OqHBGWN16zSi6tEmPwwHQ4rLPU9Plvs8/1TUBY=", + "lastModified": 1769996383, + "narHash": "sha256-AnYjnFWgS49RlqX7LrC4uA+sCCDBj0Ry/WOJ5XWAsa0=", "owner": "hercules-ci", "repo": "flake-parts", - "rev": "80daad04eddbbf5a4d883996a73f3f542fa437ac", + "rev": "57928607ea566b5db3ad13af0e57e921e6b12381", "type": "github" }, "original": { @@ -20,11 +22,11 @@ }, "nixpkgs": { "locked": { - "lastModified": 1768564909, - "narHash": "sha256-Kell/SpJYVkHWMvnhqJz/8DqQg2b6PguxVWOuadbHCc=", + "lastModified": 1770115704, + "narHash": "sha256-KHFT9UWOF2yRPlAnSXQJh6uVcgNcWlFqqiAZ7OVlHNc=", "owner": "NixOS", "repo": "nixpkgs", - "rev": "e4bae1bd10c9c57b2cf517953ab70060a828ee6f", + "rev": "e6eae2ee2110f3d31110d5c222cd395303343b08", "type": "github" }, "original": { @@ -34,21 +36,6 @@ "type": "github" } }, - "nixpkgs-lib": { - "locked": { - "lastModified": 1765674936, - "narHash": "sha256-k00uTP4JNfmejrCLJOwdObYC9jHRrr/5M/a/8L2EIdo=", - "owner": "nix-community", - "repo": "nixpkgs.lib", - "rev": "2075416fcb47225d9b68ac469a5c4801a9c4dd85", - "type": "github" - }, - "original": { - "owner": "nix-community", - "repo": "nixpkgs.lib", - "type": "github" - } - }, "root": { "inputs": { "flake-parts": "flake-parts", diff --git a/flake.nix b/flake.nix index 3807e7f..b1b0b10 100644 --- a/flake.nix +++ b/flake.nix @@ -1,6 +1,7 @@ { inputs = { flake-parts.url = "github:hercules-ci/flake-parts"; + flake-parts.inputs.nixpkgs-lib.follows = "nixpkgs"; nixpkgs.url = "github:NixOS/nixpkgs/nixos-unstable"; }; outputs = @@ -17,9 +18,9 @@ let lazuli = pkgs.buildGoModule rec { name = "lazuli"; - version = "v0.2.0"; + version = "v0.2.1"; src = pkgs.nix-gitignore.gitignoreSource [ "*.csv" "*.zip" "*.json" ] ./.; - vendorHash = "sha256-KnWoZ5UK8eigYw5uMSsLu4DIhzkSXmVHaE51Mr6hFmA="; + vendorHash = "sha256-emSCQ/WULVZ8qQ630C1bEV8hVCj2FPidws7zhDTERP4="; ldflags = [ "-X" "main.Version=${version}" diff --git a/go.mod b/go.mod index 5da1891..17fa95f 100644 --- a/go.mod +++ b/go.mod @@ -3,7 +3,7 @@ module tangled.org/karitham.dev/lazuli go 1.25.5 require ( - github.com/bluesky-social/indigo v0.0.0-20260122235001-7f2e6b43efbb + github.com/bluesky-social/indigo v0.0.0-20260202181658-ea3d39eec464 github.com/failsafe-go/failsafe-go v0.9.5 github.com/urfave/cli/v3 v3.6.2 go.etcd.io/bbolt v1.4.3 diff --git a/go.sum b/go.sum index 76221e2..ae12526 100644 --- a/go.sum +++ b/go.sum @@ -4,6 +4,8 @@ github.com/bits-and-blooms/bitset v1.24.4 h1:95H15Og1clikBrKr/DuzMXkQzECs1M6hhoG github.com/bits-and-blooms/bitset v1.24.4/go.mod h1:7hO7Gc7Pp1vODcmWvKMRA9BNmbv6a/7QIWpPxHddWR8= github.com/bluesky-social/indigo v0.0.0-20260122235001-7f2e6b43efbb h1:3FvzRkxe85/HsnQubXgdg8Vf38J5d1Sk9XmOkm2TCvY= github.com/bluesky-social/indigo v0.0.0-20260122235001-7f2e6b43efbb/go.mod h1:KIy0FgNQacp4uv2Z7xhNkV3qZiUSGuRky97s7Pa4v+o= +github.com/bluesky-social/indigo v0.0.0-20260202181658-ea3d39eec464 h1:jL6cPOk1CZ8H06sEn+WFGWufHmqkawsGyDRl+BJhQjs= +github.com/bluesky-social/indigo v0.0.0-20260202181658-ea3d39eec464/go.mod h1:VG/LeqLGNI3Ew7lsYixajnZGFfWPv144qbUddh+Oyag= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= @@ -14,6 +16,7 @@ github.com/failsafe-go/failsafe-go v0.9.5 h1:Bgt4wTKV3+n49GssB2njPZ4u5ApjvtKSIQl github.com/failsafe-go/failsafe-go v0.9.5/go.mod h1:IeRpglkcwzKagjDMh90ZhN2l4Ovt3+jemQBUbThag54= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/hashicorp/golang-lru v1.0.2 h1:dV3g9Z/unq5DpblPpw+Oqcv4dU/1omnb4Ok8iPY6p1c= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/influxdata/tdigest v0.0.1 h1:XpFptwYmnEKUqmkcDjrzffswZ3nvNeevbUSLPP/ZzIY= diff --git a/main.go b/main.go index 32dc18f..f71718f 100644 --- a/main.go +++ b/main.go @@ -10,8 +10,6 @@ import ( "io/fs" "log/slog" "os" - "os/signal" - "runtime/pprof" "slices" "strings" "time" @@ -53,7 +51,6 @@ var ( type App struct { log *slog.Logger outputFormat string - storage cache.Storage } func main() { @@ -66,28 +63,7 @@ func main() { } func run() error { - storage, err := cache.NewBoltStorage() - if err != nil { - return fmt.Errorf("open cache: %w", err) - } - - cpuFile, _ := os.Create("cpu.prof") - _ = pprof.StartCPUProfile(cpuFile) - defer pprof.StopCPUProfile() // Ensures profile is written when main exits - - c := make(chan os.Signal, 1) - signal.Notify(c, os.Interrupt) - - go func() { - <-c - fmt.Println("\nInterrupt received, saving profile and exiting...") - pprof.StopCPUProfile() - _ = cpuFile.Close() - os.Exit(0) - }() - - app := &App{storage: storage} - + app := &App{} cmd := &cli.Command{ Name: "lazuli", Usage: "Import Last.fm and Spotify listening history to Bluesky", @@ -102,9 +78,6 @@ func run() error { app.debugCommand(), app.versionCommand(), }, - After: func(ctx context.Context, cmd *cli.Command) error { - return storage.Close() - }, } return cmd.Run(context.Background(), os.Args) @@ -235,12 +208,19 @@ func (a *App) versionCommand() *cli.Command { } func (a *App) runStats(ctx context.Context, cmd *cli.Command) error { - stats, err := a.storage.Stats() + // Use read-only storage for stats to allow viewing while main process has it open + statsStorage, err := cache.NewBoltStorage(true) + if err != nil { + return fmt.Errorf("open read-only cache: %w", err) + } + defer statsStorage.Close() + + stats, err := statsStorage.Stats() if err != nil { return fmt.Errorf("failed to get database stats: %w", err) } - limiter := sync.NewRateLimiter(a.storage, 1) + limiter := sync.NewRateLimiter(statsStorage, 1) writes, global, err := limiter.Stats() if err != nil { return fmt.Errorf("failed to get rate limit stats: %w", err) @@ -289,6 +269,12 @@ func (a *App) runStats(ctx context.Context, cmd *cli.Command) error { } func (a *App) runRetry(ctx context.Context, cmd *cli.Command) error { + storage, err := cache.NewBoltStorage(false) + if err != nil { + return fmt.Errorf("open cache: %w", err) + } + defer storage.Close() + authClient, err := a.prepareAuth(ctx, cmd) if err != nil { return err @@ -296,7 +282,7 @@ func (a *App) runRetry(ctx context.Context, cmd *cli.Command) error { did := authClient.DID() dryRun := cmd.Bool("dry-run") - limiter := sync.NewRateLimiter(a.storage, 0.9) + limiter := sync.NewRateLimiter(storage, 0.9) repoClient := sync.NewRateClient(authClient.APIClient(), did, limiter) var failedRecords []struct { @@ -304,20 +290,18 @@ func (a *App) runRetry(ctx context.Context, cmd *cli.Command) error { rec sync.PlayRecord } - err = a.storage.IterateFailed(did, func(key string, rec []byte, errMsg string) error { + iterateFailed := storage.IterateFailed(did) + iterateFailed(func(key string, rec []byte, errMsg string) bool { var playRec sync.PlayRecord if err := json.Unmarshal(rec, &playRec); err != nil { - return nil + return true // continue } failedRecords = append(failedRecords, struct { key string rec sync.PlayRecord }{key, playRec}) - return nil + return true // continue }) - if err != nil { - return fmt.Errorf("failed to load failed records: %w", err) - } if len(failedRecords) == 0 { fmt.Println("No failed records to retry.") @@ -336,14 +320,14 @@ func (a *App) runRetry(ctx context.Context, cmd *cli.Command) error { continue } - res := sync.PublishBatch(ctx, repoClient, did, []*sync.PlayRecord{&fr.rec}, a.storage, sync.DefaultClientAgent) + res := sync.PublishBatch(ctx, repoClient, did, []*sync.PlayRecord{&fr.rec}, storage, sync.DefaultClientAgent) if res == nil { fmt.Printf("Successfully retried: %s - %s\n", fr.rec.ArtistName(), fr.rec.TrackName) - if err := a.storage.MarkPublished(did, fr.key); err != nil { + if err := storage.MarkPublished(did, fr.key); err != nil { a.log.Error("Failed to mark record as published", sync.ErrorAttr(err), slog.String("key", fr.key)) } - if err := a.storage.RemoveFailed(did, fr.key); err != nil { + if err := storage.RemoveFailed(did, fr.key); err != nil { a.log.Error("Failed to remove record from failed list", sync.ErrorAttr(err), slog.String("key", fr.key)) } successCount++ @@ -358,6 +342,12 @@ func (a *App) runRetry(ctx context.Context, cmd *cli.Command) error { } func (a *App) runFailed(ctx context.Context, cmd *cli.Command) error { + storage, err := cache.NewBoltStorage(true) + if err != nil { + return fmt.Errorf("open read-only cache: %w", err) + } + defer storage.Close() + authClient, err := a.prepareAuth(ctx, cmd) if err != nil { return err @@ -371,7 +361,8 @@ func (a *App) runFailed(ctx context.Context, cmd *cli.Command) error { } var failed []FailedRecord - err = a.storage.IterateFailed(did, func(key string, rec []byte, errMsg string) error { + iterateFailed := storage.IterateFailed(did) + iterateFailed(func(key string, rec []byte, errMsg string) bool { var playRec sync.PlayRecord _ = json.Unmarshal(rec, &playRec) failed = append(failed, FailedRecord{ @@ -379,11 +370,8 @@ func (a *App) runFailed(ctx context.Context, cmd *cli.Command) error { Error: errMsg, Record: playRec, }) - return nil + return true // continue }) - if err != nil { - return fmt.Errorf("failed to iterate failed records: %w", err) - } if a.outputFormat == "json" { data, _ := json.MarshalIndent(failed, "", " ") @@ -508,6 +496,12 @@ func (a *App) runExport(ctx context.Context, cmd *cli.Command) error { } func (a *App) runImport(ctx context.Context, cmd *cli.Command) error { + storage, err := cache.NewBoltStorage(false) + if err != nil { + return fmt.Errorf("open cache: %w", err) + } + defer storage.Close() + handle, password, err := a.getCredentials(cmd) if err != nil { return err @@ -525,7 +519,7 @@ func (a *App) runImport(ctx context.Context, cmd *cli.Command) error { tolerance := cmd.Duration("tolerance") if clearCache { - if err := a.storage.ClearAll(); err != nil { + if err := storage.ClearAll(); err != nil { a.log.Error("Failed to clear cache", sync.ErrorAttr(err)) } else { a.log.Info("Cache cleared") @@ -549,16 +543,16 @@ func (a *App) runImport(ctx context.Context, cmd *cli.Command) error { } a.log.Info("Authenticated", sync.DIDAttr(authClient.DID()), slog.String("pds", authClient.PDS())) - limiter := sync.NewRateLimiter(a.storage, 0.9) + limiter := sync.NewRateLimiter(storage, 0.9) repoClient := sync.NewRateClient(authClient.APIClient(), authClient.DID(), limiter) - existingRecords, err := sync.FetchExisting(ctx, repoClient, authClient.DID(), a.storage, fresh) + existingRecords, err := sync.FetchExisting(ctx, repoClient, authClient.DID(), storage, fresh) if err != nil { return fmt.Errorf("fetch existing records: %w", err) } a.log.Info("Fetched existing records", slog.Int("count", len(existingRecords))) - published, _ := a.storage.GetPublished(authClient.DID()) + published, _ := storage.GetPublished(authClient.DID()) newRecords := sync.FilterNew(records, existingRecords, published) skippedCount := len(records) - len(newRecords) a.log.Info("Filtered to new records", @@ -582,7 +576,7 @@ func (a *App) runImport(ctx context.Context, cmd *cli.Command) error { value, _ := json.Marshal(rec) newEntries[key] = value } - if err := a.storage.SaveRecords(authClient.DID(), newEntries); err != nil { + if err := storage.SaveRecords(authClient.DID(), newEntries); err != nil { return fmt.Errorf("save new records to storage: %w", err) } } @@ -592,10 +586,11 @@ func (a *App) runImport(ctx context.Context, cmd *cli.Command) error { publishOpts := sync.PublishOptions{ BatchSize: batchSize, DryRun: dryRun, + Reverse: reverse, ATProtoClient: repoClient, ProgressLog: progressLog, ClientAgent: fmt.Sprintf("lazuli/%s", Version), - Storage: a.storage, + Storage: storage, Limiter: limiter, } @@ -653,6 +648,12 @@ func (a *App) createProgressLogger() func(sync.ProgressReport) { } func (a *App) runSync(ctx context.Context, cmd *cli.Command) error { + storage, err := cache.NewBoltStorage(false) + if err != nil { + return fmt.Errorf("open cache: %w", err) + } + defer storage.Close() + authClient, err := a.prepareAuth(ctx, cmd) if err != nil { return err @@ -661,18 +662,18 @@ func (a *App) runSync(ctx context.Context, cmd *cli.Command) error { fresh := cmd.Bool("fresh") a.log.Info("Starting sync operation", sync.DIDAttr(authClient.DID()), slog.Bool("fresh", fresh)) - limiter := sync.NewRateLimiter(a.storage, 0.85) + limiter := sync.NewRateLimiter(storage, 0.85) repoClient := sync.NewRateClient(authClient.APIClient(), authClient.DID(), limiter) if fresh { - if err := a.storage.Clear(authClient.DID()); err != nil { + if err := storage.Clear(authClient.DID()); err != nil { a.log.Error("Failed to clear cache", sync.ErrorAttr(err)) } else { a.log.Info("Cache cleared") } } - existingRecords, err := sync.FetchExisting(ctx, repoClient, authClient.DID(), a.storage, fresh) + existingRecords, err := sync.FetchExisting(ctx, repoClient, authClient.DID(), storage, fresh) if err != nil { return fmt.Errorf("fetch existing records: %w", err) } @@ -683,6 +684,12 @@ func (a *App) runSync(ctx context.Context, cmd *cli.Command) error { } func (a *App) runDedupe(ctx context.Context, cmd *cli.Command) error { + storage, err := cache.NewBoltStorage(false) + if err != nil { + return fmt.Errorf("open cache: %w", err) + } + defer storage.Close() + authClient, err := a.prepareAuth(ctx, cmd) if err != nil { return fmt.Errorf("authentication failed: %w", err) @@ -696,18 +703,18 @@ func (a *App) runDedupe(ctx context.Context, cmd *cli.Command) error { slog.Bool("dry_run", dryRun), slog.Bool("fresh", fresh)) - limiter := sync.NewRateLimiter(a.storage, 0.9) + limiter := sync.NewRateLimiter(storage, 0.9) repoClient := sync.NewRateClient(authClient.APIClient(), authClient.DID(), limiter) if fresh { - if err := a.storage.Clear(authClient.DID()); err != nil { + if err := storage.Clear(authClient.DID()); err != nil { a.log.Error("Failed to clear cache", sync.ErrorAttr(err)) } else { a.log.Info("Cache cleared") } } - existingRecords, err := sync.FetchExisting(ctx, repoClient, authClient.DID(), a.storage, fresh) + existingRecords, err := sync.FetchExisting(ctx, repoClient, authClient.DID(), storage, fresh) if err != nil { return fmt.Errorf("failed to fetch existing records: %w", err) } @@ -787,7 +794,7 @@ func (a *App) runDedupe(ctx context.Context, cmd *cli.Command) error { } } - if err := a.storage.Clear(authClient.DID()); err != nil { + if err := storage.Clear(authClient.DID()); err != nil { a.log.Error("Failed to clear cache", "err", err) } diff --git a/sync/publish.go b/sync/publish.go index bad6f05..c55bbfd 100644 --- a/sync/publish.go +++ b/sync/publish.go @@ -59,9 +59,10 @@ type ( PublishOptions struct { BatchSize int DryRun bool + Reverse bool ATProtoClient ATProtoClient ProgressLog func(ProgressReport) - Storage cache.Storage + Storage cache.RecordStore Limiter RateLimiter ClientAgent string RetryDelay time.Duration @@ -85,7 +86,7 @@ type ( batchProcessor struct { Client ATProtoClient - Storage cache.Storage + Storage cache.RecordStore DID string ClientAgent string DryRun bool @@ -116,14 +117,14 @@ func NewClient(ctx context.Context, handle, password string, opts ...func(*atpro } // batchRecords iterates through storage and builds record batches -func batchRecords(ctx context.Context, storage cache.Storage, did string, batchSize int) ([]recordBatch, error) { +func batchRecords(ctx context.Context, storage cache.RecordStore, did string, batchSize int, reverse bool) ([]recordBatch, error) { var batches []recordBatch var currentBatch recordBatch - err := storage.IterateUnpublished(did, func(key string, rec []byte) error { + for key, rec := range storage.IterateUnpublished(did, reverse) { select { case <-ctx.Done(): - return ctx.Err() + return nil, ctx.Err() default: } @@ -133,23 +134,23 @@ func batchRecords(ctx context.Context, storage cache.Storage, did string, batchS if storage != nil { _ = storage.MarkFailed(did, []string{key}, "malformed record") } - return nil // Skip malformed records + continue // Skip malformed records } currentBatch.Records = append(currentBatch.Records, &record) currentBatch.Keys = append(currentBatch.Keys, key) if len(currentBatch.Records) >= batchSize { + records := make([]*PlayRecord, len(currentBatch.Records)) + copy(records, currentBatch.Records) + keys := make([]string, len(currentBatch.Keys)) + copy(keys, currentBatch.Keys) batches = append(batches, recordBatch{ - Records: append([]*PlayRecord{}, currentBatch.Records...), - Keys: append([]string{}, currentBatch.Keys...), + Records: records, + Keys: keys, }) currentBatch = recordBatch{} } - return nil - }) - if err != nil { - return nil, err } // Add the last partial batch if it has records @@ -179,7 +180,9 @@ func processBatch(ctx context.Context, batch recordBatch, processor batchProcess } } - err := failsafe.With(DefaultRetryPolicy).WithContext(ctx).Run(func() error { + policy := DefaultRetryPolicy + + err := failsafe.With(policy).WithContext(ctx).Run(func() error { return PublishBatch(ctx, processor.Client, processor.DID, batch.Records, processor.Storage, processor.ClientAgent) }) if err != nil { @@ -229,7 +232,7 @@ func Publish(ctx context.Context, client AuthClient, opts PublishOptions) publis return errorResult(startTime) } - batches, err := batchRecords(ctx, opts.Storage, client.DID(), batchSize) + batches, err := batchRecords(ctx, opts.Storage, client.DID(), batchSize, opts.Reverse) if err != nil { cancelled := ctx.Err() != nil return newPublishResult(0, 0, 0, startTime, cancelled) @@ -339,7 +342,7 @@ func logResult(success, errors int, startTime time.Time) { slog.String("rate", formatRate(ratePerMinute(success, time.Since(startTime))))) } -func PublishBatch(ctx context.Context, client ATProtoClient, did string, batch []*PlayRecord, storage cache.Storage, clientAgent string) error { +func PublishBatch(ctx context.Context, client ATProtoClient, did string, batch []*PlayRecord, storage cache.RecordStore, clientAgent string) error { if len(batch) == 0 { return nil } @@ -394,18 +397,18 @@ func FetchExisting(ctx context.Context, client RepoClient[*PlayRecord], did stri published, err := storage.GetPublished(did) if err == nil && len(published) > 0 && storage.IsValid(did) { records := make([]ExistingRecord, 0, len(published)) - err := storage.IteratePublished(did, func(key string, data []byte) error { + for _, data := range storage.IteratePublished(did, false) { var value PlayRecord if err := json.Unmarshal(data, &value); err != nil { - return nil + slog.Debug("failed to unmarshal cached record", ErrorAttr(err)) + continue } records = append(records, ExistingRecord{ URI: generateRecordURI(did, &value), Value: &value, }) - return nil - }) - if err == nil { + } + if len(records) > 0 { slog.Debug("loaded from cache", slog.Int("count", len(records))) return records, nil } diff --git a/sync/publish_test.go b/sync/publish_test.go index 41982c5..37ee409 100644 --- a/sync/publish_test.go +++ b/sync/publish_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "iter" "maps" "sync" "testing" @@ -19,7 +20,7 @@ import ( // Mock Storage type mockStorage struct { - cache.Storage + cache.RecordStore unpublished map[string][]byte published map[string]bool failed map[string]string @@ -43,25 +44,48 @@ func (m *mockStorage) SaveRecords(did string, records map[string][]byte) error { return nil } -func (m *mockStorage) IterateUnpublished(did string, fn func(key string, rec []byte) error) error { - m.mu.Lock() - keys := make([]string, 0, len(m.unpublished)) - for k := range m.unpublished { - keys = append(keys, k) +func (m *mockStorage) IterateUnpublished(did string, reverse bool) iter.Seq2[string, []byte] { + return func(yield func(key string, rec []byte) bool) { + m.mu.Lock() + keys := make([]string, 0, len(m.unpublished)) + for k := range m.unpublished { + keys = append(keys, k) + } + m.mu.Unlock() + + for _, k := range keys { + m.mu.Lock() + rec, ok := m.unpublished[k] + m.mu.Unlock() + if ok { + if !yield(k, rec) { + return + } + } + } } - m.mu.Unlock() +} - for _, k := range keys { +func (m *mockStorage) IteratePublished(did string, reverse bool) iter.Seq2[string, []byte] { + return func(yield func(key string, rec []byte) bool) { m.mu.Lock() - rec, ok := m.unpublished[k] + keys := make([]string, 0, len(m.published)) + for k := range m.published { + keys = append(keys, k) + } m.mu.Unlock() - if ok { - if err := fn(k, rec); err != nil { - return err + + for _, k := range keys { + m.mu.Lock() + rec, ok := m.unpublished[k] // Get record data from unpublished map + m.mu.Unlock() + if ok { + if !yield(k, rec) { + return + } } } } - return nil } func (m *mockStorage) MarkPublished(did string, keys ...string) error { @@ -228,7 +252,7 @@ func TestBuildRecordBatches(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - ctx := context.Background() + ctx := t.Context() if tt.ctxCancel { var cancel context.CancelFunc ctx, cancel = context.WithCancel(ctx) @@ -247,7 +271,7 @@ func TestBuildRecordBatches(t *testing.T) { storage.unpublished[fmt.Sprintf("key%d", i)] = data } - batches, err := batchRecords(ctx, storage, did, tt.batchSize) + batches, err := batchRecords(ctx, storage, did, tt.batchSize, false) if (err != nil) != tt.wantErr { t.Errorf("BuildRecordBatches() error = %v, wantErr %v", err, tt.wantErr) @@ -277,6 +301,65 @@ func TestBuildRecordBatches(t *testing.T) { } } +func TestIteratePublished(t *testing.T) { + tests := []struct { + name string + setupStorage func() *mockStorage + reverse bool + wantKeys []string + }{ + { + name: "returns only published records", + setupStorage: func() *mockStorage { + s := newMockStorage() + s.unpublished["key1"] = []byte(`{"trackName":"a"}`) + s.unpublished["key2"] = []byte(`{"trackName":"b"}`) + s.unpublished["key3"] = []byte(`{"trackName":"c"}`) + s.published["key1"] = true + s.published["key3"] = true + return s + }, + reverse: false, + wantKeys: []string{"key1", "key3"}, + }, + { + name: "handles empty published set", + setupStorage: func() *mockStorage { + s := newMockStorage() + s.unpublished["key1"] = []byte(`{"trackName":"a"}`) + return s + }, + reverse: false, + wantKeys: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + storage := tt.setupStorage() + + var gotKeys []string + for key, rec := range storage.IteratePublished("did:test", tt.reverse) { + gotKeys = append(gotKeys, key) + _ = rec + } + + if len(gotKeys) != len(tt.wantKeys) { + t.Errorf("IteratePublished() returned %d keys, want %d", len(gotKeys), len(tt.wantKeys)) + return + } + + for i, key := range gotKeys { + if key != tt.wantKeys[i] { + t.Errorf("IteratePublished() key[%d] = %s, want %s", i, key, tt.wantKeys[i]) + } + } + }) + } +} + func TestProcessBatch(t *testing.T) { tests := []struct { name string @@ -379,7 +462,7 @@ func TestProcessBatch(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - ctx := context.Background() + ctx := t.Context() result := processBatch(ctx, tt.batch, tt.processor) if result.SuccessCount != tt.wantSuccess { @@ -547,15 +630,80 @@ func TestPublish(t *testing.T) { wantSuccess: 0, wantErrors: 2, }, + { + name: "publish with reverse iteration", + opts: PublishOptions{ + BatchSize: 2, + Reverse: true, + Storage: newMockStorage(), + ClientAgent: "test-agent", + }, + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + {TrackName: "Song 3"}, + }, + wantSuccess: 3, + wantErrors: 0, + }, + { + name: "publish with empty records", + opts: PublishOptions{ + BatchSize: 2, + Storage: newMockStorage(), + ClientAgent: "test-agent", + }, + records: []PlayRecord{}, + wantSuccess: 0, + wantErrors: 0, + }, + { + name: "publish with transient errors and retry", + opts: PublishOptions{ + BatchSize: 1, + Storage: newMockStorage(), + ClientAgent: "test-agent", + }, + records: []PlayRecord{ + {TrackName: "Song 1"}, + }, + setupClient: func() *mockATProtoClient { + return &mockATProtoClient{ + applyWritesFunc: func(ctx context.Context, collection string, records []*PlayRecord) error { + return &atclient.APIError{StatusCode: 500} // Transient error + }, + } + }, + wantSuccess: 0, + wantErrors: 1, + }, + { + name: "publish with storage failure", + opts: PublishOptions{ + BatchSize: 1, + Storage: newFailingStorage(), + ClientAgent: "test-agent", + }, + records: []PlayRecord{ + {TrackName: "Song 1"}, + }, + wantSuccess: 0, + wantErrors: 1, + }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { synctest.Test(t, func(t *testing.T) { - ctx := context.Background() + ctx := t.Context() did := "did:example:123" - storage := tt.opts.Storage.(*mockStorage) + var storage *mockStorage + if fs, ok := tt.opts.Storage.(*failingStorage); ok { + storage = fs.mockStorage + } else { + storage = tt.opts.Storage.(*mockStorage) + } // Add records to storage for i, record := range tt.records { data, _ := json.Marshal(record)