diff --git a/flake.nix b/flake.nix index 6931f63..e550bb5 100644 --- a/flake.nix +++ b/flake.nix @@ -17,7 +17,7 @@ let lazuli = pkgs.buildGoModule rec { name = "lazuli"; - version = "0.1.6"; + version = "0.1.7"; src = pkgs.nix-gitignore.gitignoreSource [ "*.csv" "*.zip" "*.json" ] ./.; vendorHash = "sha256-O6R8jC8Ms5gsY2FUmuL8lTGTODfMW1CsSWuWbN27zeY="; ldflags = [ diff --git a/sync/batch_test.go b/sync/batch_test.go deleted file mode 100644 index 18d192f..0000000 --- a/sync/batch_test.go +++ /dev/null @@ -1,418 +0,0 @@ -package sync - -import ( - "context" - "encoding/json" - "errors" - "maps" - "net/http" - "strings" - "sync" - "sync/atomic" - "testing" - "testing/synctest" - "time" - - "github.com/bluesky-social/indigo/atproto/atclient" - - "tangled.org/karitham.dev/lazuli/atproto" - "tangled.org/karitham.dev/lazuli/cache" -) - -type mockRoundTripper func(req *http.Request) (*http.Response, error) - -func (f mockRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { - return f(req) -} - -// Mock Storage - -type mockStorage struct { - cache.Storage - unpublished map[string][]byte - published map[string]bool - failed map[string]string - kv map[string]int - mu sync.Mutex -} - -func newMockStorage() *mockStorage { - return &mockStorage{ - unpublished: make(map[string][]byte), - published: make(map[string]bool), - failed: make(map[string]string), - kv: make(map[string]int), - } -} - -func (m *mockStorage) SaveRecords(did string, records map[string][]byte) error { - m.mu.Lock() - defer m.mu.Unlock() - maps.Copy(m.unpublished, records) - return nil -} - -func (m *mockStorage) IterateUnpublished(did string, fn func(key string, rec []byte) error) error { - m.mu.Lock() - // Copy to avoid deadlock if fn calls back - 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 err := fn(k, rec); err != nil { - return err - } - } - } - return nil -} - -func (m *mockStorage) MarkPublished(did string, keys ...string) error { - m.mu.Lock() - defer m.mu.Unlock() - for _, k := range keys { - delete(m.unpublished, k) - m.published[k] = true - } - return nil -} - -func (m *mockStorage) MarkFailed(did string, keys []string, err string) error { - m.mu.Lock() - defer m.mu.Unlock() - for _, k := range keys { - m.failed[k] = err - } - return nil -} - -func (m *mockStorage) Get(key string) (int, error) { - m.mu.Lock() - defer m.mu.Unlock() - return m.kv[key], nil -} - -func (m *mockStorage) IncrBy(key string, n int) (int, error) { - m.mu.Lock() - defer m.mu.Unlock() - m.kv[key] += n - return m.kv[key], nil -} - -// Mock RateLimiter -type mockLimiter struct { - refunds int32 -} - -func (m *mockLimiter) AllowRead(ctx context.Context) (time.Time, error) { - return time.Now(), nil -} - -func (m *mockLimiter) AllowBulkWrite(ctx context.Context, n int) (time.Time, error) { - return time.Now(), nil -} - -func (m *mockLimiter) RefundBulkWrite(ctx context.Context, n int, chargedAt time.Time) { - atomic.AddInt32(&m.refunds, 1) -} - -func (m *mockLimiter) RefundRead(ctx context.Context, chargedAt time.Time) { - atomic.AddInt32(&m.refunds, 1) -} - -func (m *mockLimiter) Stats() (int, int, error) { - return 0, 0, nil -} - -func (m *mockLimiter) EstimatedWriteTime(n int) time.Duration { - return 0 -} - -func (m *mockLimiter) RemainingQuota() (int, int, time.Duration) { - return 10000, 35000, time.Hour -} - -// Mock ATProtoClient -type mockATProtoClient struct { - applyWritesFunc func(ctx context.Context, collection string, records []PlayRecord) error - listRecordsFunc func(ctx context.Context, collection string, limit int, cursor string) ([]atproto.RecordRef[PlayRecord], string, error) - deleteRecordFunc func(ctx context.Context, collection, rkey string) error -} - -func (m *mockATProtoClient) ApplyWrites(ctx context.Context, collection string, records []PlayRecord) error { - if m.applyWritesFunc != nil { - return m.applyWritesFunc(ctx, collection, records) - } - return nil -} - -func (m *mockATProtoClient) ListRecords(ctx context.Context, collection string, limit int, cursor string) ([]atproto.RecordRef[PlayRecord], string, error) { - if m.listRecordsFunc != nil { - return m.listRecordsFunc(ctx, collection, limit, cursor) - } - return nil, "", nil -} - -func (m *mockATProtoClient) DeleteRecord(ctx context.Context, collection, rkey string) error { - if m.deleteRecordFunc != nil { - return m.deleteRecordFunc(ctx, collection, rkey) - } - return nil -} - -// Mock AuthClient -type mockAuthClient struct { - did string -} - -func (m *mockAuthClient) APIClient() *atclient.APIClient { return nil } -func (m *mockAuthClient) DID() string { return m.did } - -type timeoutError struct{} - -func (e timeoutError) Error() string { return "timeout" } -func (e timeoutError) Timeout() bool { return true } -func (e timeoutError) Temporary() bool { return true } - -func TestIsTransientError(t *testing.T) { - tests := []struct { - name string - err error - want bool - }{ - {"nil", nil, false}, - {"generic error", errors.New("some error"), false}, - {"API 400", &atclient.APIError{StatusCode: 400}, false}, - {"API 429", &atclient.APIError{StatusCode: 429}, true}, - {"API 500", &atclient.APIError{StatusCode: 500}, true}, - {"API 503", &atclient.APIError{StatusCode: 503}, true}, - {"net timeout", timeoutError{}, true}, - {"net non-timeout", errors.New("network is down"), false}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - if got := isTransientError(tt.err); got != tt.want { - t.Errorf("isTransientError() = %v, want %v", got, tt.want) - } - }) - } -} - -func TestApplyWrites_RateClient(t *testing.T) { - ctx := context.Background() - clientAgent := "test-agent" - - tests := []struct { - name string - setupClient func() (*atproto.RateClient[PlayRecord], *mockLimiter) - records []PlayRecord - wantErr bool - wantErrMsg string - wantRefunds int32 - }{ - { - name: "empty records succeeds", - setupClient: func() (*atproto.RateClient[PlayRecord], *mockLimiter) { - limiter := &mockLimiter{} - return atproto.NewRateClient[PlayRecord](nil, "did:example:123", limiter), limiter - }, - records: nil, - wantErr: false, - wantRefunds: 0, - }, - { - name: "too many records fails", - setupClient: func() (*atproto.RateClient[PlayRecord], *mockLimiter) { - limiter := &mockLimiter{} - return atproto.NewRateClient[PlayRecord](&atclient.APIClient{}, "did:example:123", limiter), limiter - }, - records: make([]PlayRecord, 201), - wantErr: true, - wantErrMsg: "too many records in one ApplyWrites call: 201 (max 200)", - wantRefunds: 0, - }, - { - name: "transient error refunds tokens", - setupClient: func() (*atproto.RateClient[PlayRecord], *mockLimiter) { - limiter := &mockLimiter{} - apiClient := atclient.NewAPIClient("https://example.com") - apiClient.Client.Transport = mockRoundTripper(func(req *http.Request) (*http.Response, error) { - return &http.Response{ - StatusCode: 503, - Body: http.NoBody, - }, nil - }) - return atproto.NewRateClient[PlayRecord](apiClient, "did:example:123", limiter), limiter - }, - records: []PlayRecord{{TrackName: "Song 1"}}, - wantErr: true, - wantRefunds: 1, - }, - { - name: "non-transient error does NOT refund", - setupClient: func() (*atproto.RateClient[PlayRecord], *mockLimiter) { - limiter := &mockLimiter{} - apiClient := atclient.NewAPIClient("https://example.com") - apiClient.Client.Transport = mockRoundTripper(func(req *http.Request) (*http.Response, error) { - return &http.Response{ - StatusCode: 400, - Body: http.NoBody, - }, nil - }) - return atproto.NewRateClient[PlayRecord](apiClient, "did:example:123", limiter), limiter - }, - records: []PlayRecord{{TrackName: "Song 1"}}, - wantErr: true, - wantRefunds: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - client, limiter := tt.setupClient() - err := client.ApplyWrites(ctx, "test", tt.records) - - if tt.wantErr { - if err == nil { - t.Error("expected error, got nil") - } else if tt.wantErrMsg != "" && err.Error() != tt.wantErrMsg { - t.Errorf("error msg = %q, want %q", err.Error(), tt.wantErrMsg) - } - } else if err != nil { - t.Errorf("unexpected error: %v", err) - } - - if got := atomic.LoadInt32(&limiter.refunds); got != tt.wantRefunds { - t.Errorf("refunds = %d, want %d", got, tt.wantRefunds) - } - }) - } - _ = clientAgent -} - -func TestPublishBatch(t *testing.T) { - ctx := context.Background() - did := "did:example:123" - batch := []PlayRecord{{TrackName: "Song 1"}} - clientAgent := "test-agent" - - t.Run("Success", func(t *testing.T) { - storage := newMockStorage() - client := &mockATProtoClient{} - err := PublishBatch(ctx, client, did, batch, storage, clientAgent) - if err != nil { - t.Fatal(err) - } - if len(storage.published) != 1 { - t.Errorf("expected 1 published record, got %d", len(storage.published)) - } - }) - - t.Run("ApplyWrites failure", func(t *testing.T) { - storage := newMockStorage() - expectedErr := errors.New("apply failed") - client := &mockATProtoClient{ - applyWritesFunc: func(ctx context.Context, collection string, records []PlayRecord) error { - return expectedErr - }, - } - err := PublishBatch(ctx, client, did, batch, storage, clientAgent) - if !errors.Is(err, expectedErr) { - t.Errorf("expected error %v, got %v", expectedErr, err) - } - if len(storage.published) != 0 { - t.Error("expected 0 published records") - } - }) - - t.Run("Storage failure after ApplyWrites success", func(t *testing.T) { - storage := &failingStorage{} - client := &mockATProtoClient{} - err := PublishBatch(ctx, client, did, batch, storage, clientAgent) - if err == nil || !strings.Contains(err.Error(), "failed to save records") { - t.Errorf("expected storage save error, got %v", err) - } - }) -} - -type failingStorage struct { - mockStorage -} - -func (s *failingStorage) SaveRecords(did string, records map[string][]byte) error { - return errors.New("failed to save records") -} - -func TestPublish_Iterative(t *testing.T) { - ctx := context.Background() - did := "did:example:123" - clientAgent := "test-agent" - - rec1, _ := json.Marshal(PlayRecord{TrackName: "Song 1"}) - rec2, _ := json.Marshal(PlayRecord{TrackName: "Song 2"}) - - t.Run("Retry on transient error", func(t *testing.T) { - synctest.Test(t, func(t *testing.T) { - storage := newMockStorage() - storage.SaveRecords(did, map[string][]byte{"k1": rec1, "k2": rec2}) - - var attempts int32 - client := &mockATProtoClient{ - applyWritesFunc: func(ctx context.Context, collection string, records []PlayRecord) error { - if atomic.AddInt32(&attempts, 1) <= 2 { - return &atclient.APIError{StatusCode: 503} - } - return nil - }, - } - - res := Publish(ctx, &mockAuthClient{did: did}, PublishOptions{ - BatchSize: 1, - ATProtoClient: client, - Storage: storage, - ClientAgent: clientAgent, - }) - - if res.SuccessCount != 2 { - t.Errorf("expected 2 successes, got %d", res.SuccessCount) - } - if atomic.LoadInt32(&attempts) < 3 { - t.Errorf("expected at least 3 attempts (2 fails + 1 success), got %d", attempts) - } - }) - }) - - t.Run("Fail fast on non-transient error", func(t *testing.T) { - storage := newMockStorage() - storage.SaveRecords(did, map[string][]byte{"k1": rec1}) - - client := &mockATProtoClient{ - applyWritesFunc: func(ctx context.Context, collection string, records []PlayRecord) error { - return &atclient.APIError{StatusCode: 400} - }, - } - - res := Publish(ctx, &mockAuthClient{did: did}, PublishOptions{ - BatchSize: 1, - ATProtoClient: client, - Storage: storage, - ClientAgent: clientAgent, - }) - - if res.SuccessCount != 0 { - t.Errorf("expected 0 successes, got %d", res.SuccessCount) - } - if res.ErrorCount != 1 { - t.Errorf("expected 1 error, got %d", res.ErrorCount) - } - }) -} diff --git a/sync/publish.go b/sync/publish.go index 30e1875..aa70d98 100644 --- a/sync/publish.go +++ b/sync/publish.go @@ -34,6 +34,20 @@ const ( BaseRetryDelay = 2 * time.Second ) +var DefaultRetryPolicy = retrypolicy.NewBuilder[struct{}](). + WithMaxRetries(10). + WithBackoff(BaseRetryDelay, 5*time.Minute). + HandleIf(func(_ struct{}, err error) bool { + return atproto.IsTransientError(err) + }). + OnRetryScheduled(func(e failsafe.ExecutionScheduledEvent[struct{}]) { + slog.Warn("batch failed with transient error, retrying", + slog.Duration("retryDelay", e.Delay), + ErrorAttr(e.LastError()), + slog.Int("attempt", e.Attempts())) + }). + Build() + type ( ATProtoClient = atproto.RepoClient[PlayRecord] AuthClient = atproto.AuthClient @@ -53,7 +67,7 @@ type ( RetryDelay time.Duration } - PublishResult struct { + publishResult struct { SuccessCount int `json:"successCount"` ErrorCount int `json:"errorCount"` Cancelled bool `json:"cancelled"` @@ -63,6 +77,26 @@ type ( FirstRecordTime time.Time `json:"firstRecordTime"` LastRecordTime time.Time `json:"lastRecordTime"` } + + recordBatch struct { + Records []PlayRecord + Keys []string + } + + batchProcessor struct { + Client ATProtoClient + Storage cache.Storage + DID string + ClientAgent string + DryRun bool + } + + batchResult struct { + SuccessCount int + ErrorCount int + Duration time.Duration + Errors []error + } ) func NewRateLimiter(kv atproto.KVStore, maxPercent float32) RateLimiter { @@ -81,174 +115,187 @@ func NewClient(ctx context.Context, handle, password string, opts ...func(*atpro return atproto.NewClient(ctx, handle, password, opts...) } -func Publish(ctx context.Context, client AuthClient, opts PublishOptions) PublishResult { - startTime := time.Now() +// batchRecords iterates through storage and builds record batches +func batchRecords(ctx context.Context, storage cache.Storage, did string, batchSize int) ([]recordBatch, error) { + var batches []recordBatch + var currentBatch recordBatch - retryDelay := cmp.Or(opts.RetryDelay, BaseRetryDelay) - batchSize := cmp.Or(opts.BatchSize, DefaultBatchSize) + err := storage.IterateUnpublished(did, func(key string, rec []byte) error { + select { + case <-ctx.Done(): + return ctx.Err() + default: + } - atprotoClient, err := atproto.BuildClient(client, opts.ATProtoClient) - if err != nil { - return PublishResult{ - SuccessCount: 0, - ErrorCount: 0, - Cancelled: false, - Duration: time.Since(startTime), - TotalRecords: 0, + var record PlayRecord + if err := json.Unmarshal(rec, &record); err != nil { + slog.Error("malformed record in storage", slog.String("key", key), ErrorAttr(err)) + if storage != nil { + _ = storage.MarkFailed(did, []string{key}, "malformed record") + } + return nil // Skip malformed records } - } - totalRecords := 0 - if opts.Storage != nil { - _ = opts.Storage.IterateUnpublished(client.DID(), func(key string, rec []byte) error { - totalRecords++ - return nil - }) + currentBatch.Records = append(currentBatch.Records, record) + currentBatch.Keys = append(currentBatch.Keys, key) + + if len(currentBatch.Records) >= batchSize { + batches = append(batches, recordBatch{ + Records: append([]PlayRecord{}, currentBatch.Records...), + Keys: append([]string{}, currentBatch.Keys...), + }) + currentBatch = recordBatch{} + } + return nil + }) + if err != nil { + return nil, err } - if totalRecords == 0 { - return PublishResult{} + // Add the last partial batch if it has records + if len(currentBatch.Records) > 0 { + batches = append(batches, currentBatch) } - slog.Info("starting iterative import", - slog.Int("total_records", totalRecords), - slog.Int("batch_size", batchSize), - slog.Int("daily_write_limit", atproto.WriteLimitDay), - slog.Int("daily_token_limit", atproto.GlobalLimitDay), - slog.String("rate_limit", fmt.Sprintf("1 write per %.1fs", 86400.0/atproto.WriteLimitDay))) + return batches, nil +} - tracker := NewProgressTracker(totalRecords, opts.Limiter) - progressLog := defaultProgressLog(opts.ProgressLog) - totalSuccess := 0 - totalErrors := 0 +// processBatch processes a single batch of records with retries +func processBatch(ctx context.Context, batch recordBatch, processor batchProcessor) batchResult { + if len(batch.Records) == 0 { + return batchResult{} + } - var batch []PlayRecord - var batchKeys []string + start := time.Now() - processBatch := func() error { - if len(batch) == 0 { - return nil + if processor.DryRun { + for _, r := range batch.Records { + tid := syntax.NewTIDFromTime(r.PlayedTime.Time, 0) + slog.Info("would publish record (dry run)", trackAttr(r), slog.String("rkey", string(tid))) } + return batchResult{ + SuccessCount: len(batch.Records), + Duration: time.Since(start), + } + } + + err := failsafe.With(DefaultRetryPolicy).WithContext(ctx).Run(func() error { + return PublishBatch(ctx, processor.Client, processor.DID, batch.Records, processor.Storage, processor.ClientAgent) + }) + if err != nil { + slog.Error("batch failed after retries", + ErrorAttr(err), + slog.Int("count", len(batch.Records))) - if opts.DryRun { - for _, r := range batch { - tid := syntax.NewTIDFromTime(r.PlayedTime.Time, 0) - slog.Info("would publish record (dry run)", trackAttr(r), slog.String("rkey", string(tid))) + if processor.Storage != nil { + if markErr := processor.Storage.MarkFailed(processor.DID, batch.Keys, err.Error()); markErr != nil { + slog.Error("failed to mark records as failed", ErrorAttr(markErr)) } - totalSuccess += len(batch) - tracker.Increment(len(batch)) - slog.Debug("batch dry run completed", - slog.Int("count", len(batch)), - slog.Int("completed", tracker.Completed)) - batch = batch[:0] - batchKeys = batchKeys[:0] - return nil } - slog.Debug("processing batch", - slog.Int("count", len(batch)), - slog.Int("completed", tracker.Completed), - slog.Int("total", tracker.Total)) - - did := client.DID() - retryPolicy := retrypolicy.NewBuilder[any](). - WithMaxRetries(10). - WithBackoff(retryDelay, 5*time.Minute). - HandleIf(func(_ any, err error) bool { - return isTransientError(err) - }). - OnRetryScheduled(func(e failsafe.ExecutionScheduledEvent[any]) { - slog.Warn("batch failed with transient error, retrying", - slog.Int("count", len(batch)), - slog.Duration("retryDelay", e.Delay), - ErrorAttr(e.LastError()), - slog.Int("attempt", e.Attempts())) - }). - Build() - - err := failsafe.With(retryPolicy).WithContext(ctx).Run(func() error { - return PublishBatch(ctx, atprotoClient, did, batch, opts.Storage, opts.ClientAgent) - }) - if err != nil { - slog.Error("batch failed after retries", - ErrorAttr(err), - slog.Int("count", len(batch))) + return batchResult{ + ErrorCount: len(batch.Records), + Duration: time.Since(start), + Errors: []error{err}, + } + } - if opts.Storage != nil { - if markErr := opts.Storage.MarkFailed(did, batchKeys, err.Error()); markErr != nil { - slog.Error("failed to mark records as failed", ErrorAttr(markErr)) - } - } + return batchResult{ + SuccessCount: len(batch.Records), + Duration: time.Since(start), + } +} - totalErrors += len(batch) - tracker.IncrementErrors(len(batch)) +// aggregate combines batch results into final publish result +func aggregate(results []batchResult, startTime time.Time) publishResult { + totalSuccess := 0 + totalErrors := 0 - batch = batch[:0] - batchKeys = batchKeys[:0] - return nil - } + for _, result := range results { + totalSuccess += result.SuccessCount + totalErrors += result.ErrorCount + } - totalSuccess += len(batch) - tracker.Increment(len(batch)) - slog.Debug("batch published", - slog.Int("count", len(batch)), - slog.Int("completed", tracker.Completed), - slog.Int("total", tracker.Total)) + logResult(totalSuccess, totalErrors, startTime) + return newPublishResult(totalSuccess, totalErrors, totalSuccess+totalErrors, startTime, false) +} - if tracker.ShouldLog() { - progressLog(tracker.Report()) - } +func Publish(ctx context.Context, client AuthClient, opts PublishOptions) publishResult { + startTime := time.Now() + batchSize := cmp.Or(opts.BatchSize, DefaultBatchSize) - batch = batch[:0] - batchKeys = batchKeys[:0] - return nil + atprotoClient, err := atproto.BuildClient(client, opts.ATProtoClient) + if err != nil { + return errorResult(startTime) + } + + batches, err := batchRecords(ctx, opts.Storage, client.DID(), batchSize) + if err != nil { + cancelled := ctx.Err() != nil + return newPublishResult(0, 0, 0, startTime, cancelled) + } + + if len(batches) == 0 { + return publishResult{} + } + + slog.Info("starting iterative import", + slog.Int("total_records", countTotalRecords(batches)), + slog.Int("batch_size", batchSize), + slog.Int("daily_write_limit", atproto.WriteLimitDay), + slog.Int("daily_token_limit", atproto.GlobalLimitDay), + slog.String("rate_limit", fmt.Sprintf("1 write per %.1fs", 86400.0/atproto.WriteLimitDay))) + + tracker := NewProgressTracker(countTotalRecords(batches), opts.Limiter) + progressLog := defaultProgressLog(opts.ProgressLog) + + processor := batchProcessor{ + Client: atprotoClient, + Storage: opts.Storage, + DID: client.DID(), + ClientAgent: opts.ClientAgent, + DryRun: opts.DryRun, } - err = opts.Storage.IterateUnpublished(client.DID(), func(key string, rec []byte) error { + var results []batchResult + for _, batch := range batches { select { case <-ctx.Done(): - return ctx.Err() + return aggregate(results, startTime) default: } - var record PlayRecord - if err := json.Unmarshal(rec, &record); err != nil { - slog.Error("malformed record in storage", slog.String("key", key), ErrorAttr(err)) - if opts.Storage != nil { - _ = opts.Storage.MarkFailed(client.DID(), []string{key}, "malformed record") - } - totalErrors++ - tracker.IncrementErrors(1) - return nil - } + result := processBatch(ctx, batch, processor) + results = append(results, result) - batch = append(batch, record) - batchKeys = append(batchKeys, key) + // Update progress tracking + tracker.Increment(result.SuccessCount + result.ErrorCount) + tracker.IncrementErrors(result.ErrorCount) - if len(batch) >= batchSize { - if err := processBatch(); err != nil { - return err - } + if tracker.ShouldLog() { + progressLog(tracker.Report()) } - return nil - }) - - if err == nil && len(batch) > 0 { - err = processBatch() } - cancelled := false - if err != nil { - slog.Error("import interrupted", ErrorAttr(err)) - cancelled = true - } + return aggregate(results, startTime) +} - logResult(totalSuccess, totalErrors, startTime) - return newPublishResult(totalSuccess, totalErrors, totalRecords, startTime, cancelled) +func countTotalRecords(batches []recordBatch) int { + total := 0 + for _, batch := range batches { + total += len(batch.Records) + } + return total } -func isTransientError(err error) bool { - return atproto.IsTransientError(err) +func errorResult(startTime time.Time) publishResult { + return publishResult{ + SuccessCount: 0, + ErrorCount: 0, + Cancelled: false, + Duration: time.Since(startTime), + TotalRecords: 0, + } } func defaultProgressLog(f func(ProgressReport)) func(ProgressReport) { @@ -268,8 +315,8 @@ func defaultProgressLog(f func(ProgressReport)) func(ProgressReport) { } } -func newPublishResult(success, errors, total int, start time.Time, cancelled bool) PublishResult { - return PublishResult{ +func newPublishResult(success, errors, total int, start time.Time, cancelled bool) publishResult { + return publishResult{ SuccessCount: success, ErrorCount: errors, Cancelled: cancelled, @@ -384,7 +431,7 @@ func fetchExistingLoop(ctx context.Context, client RepoClient[PlayRecord], did s cursor string } - retryPolicy := retrypolicy.NewBuilder[fetchResult](). + fetchRetryPolicy := retrypolicy.NewBuilder[fetchResult](). WithMaxRetries(10). WithBackoff(BaseRetryDelay, 5*time.Minute). HandleIf(func(_ fetchResult, err error) bool { @@ -405,7 +452,7 @@ func fetchExistingLoop(ctx context.Context, client RepoClient[PlayRecord], did s default: } - result, err := failsafe.With(retryPolicy). + result, err := failsafe.With(fetchRetryPolicy). WithContext(ctx). Get(func() (fetchResult, error) { recs, next, err := client.ListRecords(ctx, RecordType, batchSize, cursor) diff --git a/sync/publish_test.go b/sync/publish_test.go new file mode 100644 index 0000000..024ab83 --- /dev/null +++ b/sync/publish_test.go @@ -0,0 +1,614 @@ +package sync + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "maps" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/bluesky-social/indigo/atproto/atclient" + + "tangled.org/karitham.dev/lazuli/atproto" + "tangled.org/karitham.dev/lazuli/cache" +) + +// Mock Storage +type mockStorage struct { + cache.Storage + unpublished map[string][]byte + published map[string]bool + failed map[string]string + kv map[string]int + mu sync.Mutex +} + +func newMockStorage() *mockStorage { + return &mockStorage{ + unpublished: make(map[string][]byte), + published: make(map[string]bool), + failed: make(map[string]string), + kv: make(map[string]int), + } +} + +func (m *mockStorage) SaveRecords(did string, records map[string][]byte) error { + m.mu.Lock() + defer m.mu.Unlock() + maps.Copy(m.unpublished, records) + 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) + } + m.mu.Unlock() + + for _, k := range keys { + m.mu.Lock() + rec, ok := m.unpublished[k] + m.mu.Unlock() + if ok { + if err := fn(k, rec); err != nil { + return err + } + } + } + return nil +} + +func (m *mockStorage) MarkPublished(did string, keys ...string) error { + m.mu.Lock() + defer m.mu.Unlock() + for _, k := range keys { + delete(m.unpublished, k) + m.published[k] = true + } + return nil +} + +func (m *mockStorage) MarkFailed(did string, keys []string, err string) error { + m.mu.Lock() + defer m.mu.Unlock() + for _, k := range keys { + m.failed[k] = err + } + return nil +} + +func (m *mockStorage) Get(key string) (int, error) { + m.mu.Lock() + defer m.mu.Unlock() + return m.kv[key], nil +} + +func (m *mockStorage) IncrBy(key string, n int) (int, error) { + m.mu.Lock() + defer m.mu.Unlock() + m.kv[key] += n + return m.kv[key], nil +} + +// Mock ATProtoClient +type mockATProtoClient struct { + applyWritesFunc func(ctx context.Context, collection string, records []PlayRecord) error + listRecordsFunc func(ctx context.Context, collection string, limit int, cursor string) ([]atproto.RecordRef[PlayRecord], string, error) + deleteRecordFunc func(ctx context.Context, collection, rkey string) error +} + +func (m *mockATProtoClient) ApplyWrites(ctx context.Context, collection string, records []PlayRecord) error { + if m.applyWritesFunc != nil { + return m.applyWritesFunc(ctx, collection, records) + } + return nil +} + +func (m *mockATProtoClient) ListRecords(ctx context.Context, collection string, limit int, cursor string) ([]atproto.RecordRef[PlayRecord], string, error) { + if m.listRecordsFunc != nil { + return m.listRecordsFunc(ctx, collection, limit, cursor) + } + return nil, "", nil +} + +func (m *mockATProtoClient) DeleteRecord(ctx context.Context, collection, rkey string) error { + if m.deleteRecordFunc != nil { + return m.deleteRecordFunc(ctx, collection, rkey) + } + return nil +} + +// Mock AuthClient +type mockAuthClient struct { + did string +} + +func (m *mockAuthClient) APIClient() *atclient.APIClient { return nil } +func (m *mockAuthClient) DID() string { return m.did } + +type timeoutError struct{} + +func (e timeoutError) Error() string { return "timeout" } +func (e timeoutError) Timeout() bool { return true } +func (e timeoutError) Temporary() bool { return true } + +type failingStorage struct { + *mockStorage +} + +func newFailingStorage() *failingStorage { + return &failingStorage{ + mockStorage: newMockStorage(), + } +} + +func (s *failingStorage) SaveRecords(did string, records map[string][]byte) error { + return errors.New("failed to save records") +} + +func TestBuildRecordBatches(t *testing.T) { + tests := []struct { + name string + records []PlayRecord + batchSize int + wantBatches int + wantErr bool + ctxCancel bool + }{ + { + name: "empty storage", + records: []PlayRecord{}, + batchSize: 2, + wantBatches: 0, + wantErr: false, + }, + { + name: "single batch", + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + }, + batchSize: 5, + wantBatches: 1, + wantErr: false, + }, + { + name: "multiple exact batches", + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + {TrackName: "Song 3"}, + {TrackName: "Song 4"}, + }, + batchSize: 2, + wantBatches: 2, + wantErr: false, + }, + { + name: "partial final batch", + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + {TrackName: "Song 3"}, + }, + batchSize: 2, + wantBatches: 2, + wantErr: false, + }, + { + name: "context cancelled", + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + }, + batchSize: 2, + wantBatches: 0, + wantErr: true, + ctxCancel: true, + }, + { + name: "malformed records skipped", + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + }, + batchSize: 2, + wantBatches: 1, + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + ctx := context.Background() + if tt.ctxCancel { + var cancel context.CancelFunc + ctx, cancel = context.WithCancel(ctx) + cancel() + } + + storage := newMockStorage() + did := "did:example:123" + + // Add records to storage + for i, record := range tt.records { + data, _ := json.Marshal(record) + if tt.name == "malformed records skipped" && i == 1 { + data = []byte("invalid json") + } + storage.unpublished[fmt.Sprintf("key%d", i)] = data + } + + batches, err := batchRecords(ctx, storage, did, tt.batchSize) + + if (err != nil) != tt.wantErr { + t.Errorf("BuildRecordBatches() error = %v, wantErr %v", err, tt.wantErr) + return + } + + if len(batches) != tt.wantBatches { + t.Errorf("BuildRecordBatches() batches = %d, want %d", len(batches), tt.wantBatches) + } + + if !tt.wantErr { + totalRecords := 0 + for _, batch := range batches { + totalRecords += len(batch.Records) + } + + expectedRecords := len(tt.records) + if tt.name == "malformed records skipped" { + expectedRecords = len(tt.records) - 1 // Skip malformed record + } + + if totalRecords != expectedRecords { + t.Errorf("BuildRecordBatches() total records = %d, want %d", totalRecords, expectedRecords) + } + } + }) + } +} + +func TestProcessBatch(t *testing.T) { + tests := []struct { + name string + batch recordBatch + processor batchProcessor + wantSuccess int + wantError int + wantErr bool + setupClient func() *mockATProtoClient + setupStorage func() cache.Storage + }{ + { + name: "empty batch", + batch: recordBatch{ + Records: []PlayRecord{}, + Keys: []string{}, + }, + processor: batchProcessor{ + Client: &mockATProtoClient{}, + Storage: newMockStorage(), + }, + wantSuccess: 0, + wantError: 0, + wantErr: false, + }, + { + name: "successful batch", + batch: recordBatch{ + Records: []PlayRecord{{TrackName: "Song 1"}, {TrackName: "Song 2"}}, + Keys: []string{"key1", "key2"}, + }, + processor: batchProcessor{ + Client: &mockATProtoClient{}, + Storage: newMockStorage(), + DID: "did:example:123", + ClientAgent: "test-agent", + }, + wantSuccess: 2, + wantError: 0, + wantErr: false, + }, + { + name: "dry run batch", + batch: recordBatch{ + Records: []PlayRecord{{TrackName: "Song 1"}, {TrackName: "Song 2"}}, + Keys: []string{"key1", "key2"}, + }, + processor: batchProcessor{ + Client: &mockATProtoClient{}, + Storage: newMockStorage(), + DID: "did:example:123", + ClientAgent: "test-agent", + DryRun: true, + }, + wantSuccess: 2, + wantError: 0, + wantErr: false, + }, + { + name: "batch with apply writes failure", + batch: recordBatch{ + Records: []PlayRecord{{TrackName: "Song 1"}}, + Keys: []string{"key1"}, + }, + processor: batchProcessor{ + Client: func() *mockATProtoClient { + return &mockATProtoClient{ + applyWritesFunc: func(ctx context.Context, collection string, records []PlayRecord) error { + return errors.New("apply writes failed") + }, + } + }(), + Storage: newMockStorage(), + DID: "did:example:123", + ClientAgent: "test-agent", + }, + wantSuccess: 0, + wantError: 1, + wantErr: false, + }, + { + name: "batch with storage failure", + batch: recordBatch{ + Records: []PlayRecord{{TrackName: "Song 1"}}, + Keys: []string{"key1"}, + }, + processor: batchProcessor{ + Client: &mockATProtoClient{}, + Storage: newFailingStorage(), + DID: "did:example:123", + ClientAgent: "test-agent", + }, + wantSuccess: 0, + wantError: 1, + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + ctx := context.Background() + result := processBatch(ctx, tt.batch, tt.processor) + + if result.SuccessCount != tt.wantSuccess { + t.Errorf("ProcessBatch() success count = %d, want %d", result.SuccessCount, tt.wantSuccess) + } + + if result.ErrorCount != tt.wantError { + t.Errorf("ProcessBatch() error count = %d, want %d", result.ErrorCount, tt.wantError) + } + + if tt.wantError > 0 && len(result.Errors) == 0 { + t.Error("ProcessBatch() expected errors but got none") + } + }) + } +} + +func TestAggregateResults(t *testing.T) { + tests := []struct { + name string + results []batchResult + startTime time.Time + wantSuccess int + wantErrors int + wantTotal int + wantDuration bool + wantRatePerMin bool + }{ + { + name: "empty results", + results: []batchResult{}, + wantSuccess: 0, + wantErrors: 0, + wantTotal: 0, + }, + { + name: "single successful result", + results: []batchResult{ + {SuccessCount: 5, ErrorCount: 0, Duration: time.Second}, + }, + wantSuccess: 5, + wantErrors: 0, + wantTotal: 5, + wantDuration: true, + wantRatePerMin: true, + }, + { + name: "multiple mixed results", + results: []batchResult{ + {SuccessCount: 3, ErrorCount: 1, Duration: time.Second}, + {SuccessCount: 2, ErrorCount: 0, Duration: time.Second}, + {SuccessCount: 0, ErrorCount: 2, Duration: time.Second}, + }, + wantSuccess: 5, + wantErrors: 3, + wantTotal: 8, + wantDuration: true, + wantRatePerMin: true, + }, + { + name: "all errors", + results: []batchResult{ + {SuccessCount: 0, ErrorCount: 3, Duration: time.Second}, + {SuccessCount: 0, ErrorCount: 2, Duration: time.Second}, + }, + wantSuccess: 0, + wantErrors: 5, + wantTotal: 5, + wantDuration: true, + wantRatePerMin: false, // 0 success rate + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + startTime := time.Now() + if !tt.startTime.IsZero() { + startTime = tt.startTime + } + + result := aggregate(tt.results, startTime) + + if result.SuccessCount != tt.wantSuccess { + t.Errorf("AggregateResults() success = %d, want %d", result.SuccessCount, tt.wantSuccess) + } + + if result.ErrorCount != tt.wantErrors { + t.Errorf("AggregateResults() errors = %d, want %d", result.ErrorCount, tt.wantErrors) + } + + if result.TotalRecords != tt.wantTotal { + t.Errorf("AggregateResults() total = %d, want %d", result.TotalRecords, tt.wantTotal) + } + + if tt.wantDuration && result.Duration == 0 { + t.Error("AggregateResults() expected non-zero duration") + } + + if tt.wantRatePerMin && result.RecordsPerMinute == 0 && tt.wantSuccess > 0 { + t.Error("AggregateResults() expected non-zero rate per minute") + } + }) + } +} + +func TestPublish(t *testing.T) { + tests := []struct { + name string + opts PublishOptions + records []PlayRecord + setupClient func() *mockATProtoClient + wantSuccess int + wantErrors int + wantCancelled bool + }{ + { + name: "successful publish", + opts: PublishOptions{ + BatchSize: 2, + Storage: newMockStorage(), + ClientAgent: "test-agent", + }, + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + }, + wantSuccess: 2, + wantErrors: 0, + }, + { + name: "dry run publish", + opts: PublishOptions{ + BatchSize: 2, + DryRun: true, + Storage: newMockStorage(), + ClientAgent: "test-agent", + }, + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + }, + wantSuccess: 2, + wantErrors: 0, + }, + { + name: "publish with client errors", + opts: PublishOptions{ + BatchSize: 1, + Storage: newMockStorage(), + ClientAgent: "test-agent", + }, + records: []PlayRecord{ + {TrackName: "Song 1"}, + {TrackName: "Song 2"}, + }, + setupClient: func() *mockATProtoClient { + return &mockATProtoClient{ + applyWritesFunc: func(ctx context.Context, collection string, records []PlayRecord) error { + return &atclient.APIError{StatusCode: 400} // Non-transient error + }, + } + }, + wantSuccess: 0, + wantErrors: 2, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + ctx := context.Background() + did := "did:example:123" + + storage := tt.opts.Storage.(*mockStorage) + // Add records to storage + for i, record := range tt.records { + data, _ := json.Marshal(record) + storage.unpublished[fmt.Sprintf("key%d", i)] = data + } + + client := &mockAuthClient{did: did} + if tt.setupClient != nil { + tt.opts.ATProtoClient = tt.setupClient() + } else { + tt.opts.ATProtoClient = &mockATProtoClient{} + } + + result := Publish(ctx, client, tt.opts) + + if result.SuccessCount != tt.wantSuccess { + t.Errorf("Publish() success = %d, want %d", result.SuccessCount, tt.wantSuccess) + } + + if result.ErrorCount != tt.wantErrors { + t.Errorf("Publish() errors = %d, want %d", result.ErrorCount, tt.wantErrors) + } + + if result.Cancelled != tt.wantCancelled { + t.Errorf("Publish() cancelled = %v, want %v", result.Cancelled, tt.wantCancelled) + } + }) + }) + } +} + +func TestIsTransientError(t *testing.T) { + tests := []struct { + name string + err error + want bool + }{ + {"nil", nil, false}, + {"generic error", errors.New("some error"), false}, + {"API 400", &atclient.APIError{StatusCode: 400}, false}, + {"API 429", &atclient.APIError{StatusCode: 429}, true}, + {"API 500", &atclient.APIError{StatusCode: 500}, true}, + {"API 503", &atclient.APIError{StatusCode: 503}, true}, + {"net timeout", timeoutError{}, true}, + {"net non-timeout", errors.New("network is down"), false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if got := IsTransientError(tt.err); got != tt.want { + t.Errorf("IsTransientError() = %v, want %v", got, tt.want) + } + }) + } +}