From 575b439c3c4f76d29872e90ff04040275782c3a7 Mon Sep 17 00:00:00 2001 From: karitham Date: Thu, 22 Jan 2026 14:07:03 +0100 Subject: [PATCH] sync/rate: respect burst limit, code quality --- cache/bbolt.go | 85 ++++++++-- cache/cache_test.go | 17 -- cache/storage.go | 4 +- flake.nix | 4 +- go.mod | 31 ++-- go.sum | 89 ++++++---- main.go | 46 ++++-- sync/adapter.go | 389 ++++++++++++++++++++++++++++++++++++-------- sync/batch_test.go | 364 +++++++++++++++++++++++++++++++++++++++++ sync/config.go | 3 +- sync/import_test.go | 43 ++++- sync/progress.go | 42 +++-- sync/publish.go | 315 +++++------------------------------ sync/rate.go | 296 +++++++++++++++++++++++---------- sync/rate_test.go | 160 +++++++++++++++--- 15 files changed, 1331 insertions(+), 557 deletions(-) create mode 100644 sync/batch_test.go diff --git a/cache/bbolt.go b/cache/bbolt.go index bd07833..619aeb7 100644 --- a/cache/bbolt.go +++ b/cache/bbolt.go @@ -10,6 +10,7 @@ import ( "time" "go.etcd.io/bbolt" + berr "go.etcd.io/bbolt/errors" ) var _ Storage = (*BoltStorage)(nil) @@ -244,13 +245,13 @@ func (s *BoltStorage) Clear(did string) error { } return s.db.Update(func(tx *bbolt.Tx) error { - if err := tx.DeleteBucket([]byte(recordsBucket(did))); err != nil && !errors.Is(err, bbolt.ErrBucketNotFound) { + if err := tx.DeleteBucket([]byte(recordsBucket(did))); err != nil && !errors.Is(err, berr.ErrBucketNotFound) { return fmt.Errorf("delete records bucket: %w", err) } - if err := tx.DeleteBucket([]byte(processedBucket(did))); err != nil && !errors.Is(err, bbolt.ErrBucketNotFound) { + if err := tx.DeleteBucket([]byte(processedBucket(did))); err != nil && !errors.Is(err, berr.ErrBucketNotFound) { return fmt.Errorf("delete processed bucket: %w", err) } - if err := tx.DeleteBucket([]byte(failedBucket(did))); err != nil && !errors.Is(err, bbolt.ErrBucketNotFound) { + if err := tx.DeleteBucket([]byte(failedBucket(did))); err != nil && !errors.Is(err, berr.ErrBucketNotFound) { return fmt.Errorf("delete failed bucket: %w", err) } metaBkt := tx.Bucket([]byte(metaBucket())) @@ -278,22 +279,22 @@ func (s *BoltStorage) ClearAll() error { }) for _, did := range dids { - if err := tx.DeleteBucket([]byte(recordsBucket(did))); err != nil && !errors.Is(err, bbolt.ErrBucketNotFound) { + if err := tx.DeleteBucket([]byte(recordsBucket(did))); err != nil && !errors.Is(err, berr.ErrBucketNotFound) { return fmt.Errorf("delete records for %s: %w", did, err) } - if err := tx.DeleteBucket([]byte(processedBucket(did))); err != nil && !errors.Is(err, bbolt.ErrBucketNotFound) { + if err := tx.DeleteBucket([]byte(processedBucket(did))); err != nil && !errors.Is(err, berr.ErrBucketNotFound) { return fmt.Errorf("delete processed for %s: %w", did, err) } - if err := tx.DeleteBucket([]byte(failedBucket(did))); err != nil && !errors.Is(err, bbolt.ErrBucketNotFound) { + if err := tx.DeleteBucket([]byte(failedBucket(did))); err != nil && !errors.Is(err, berr.ErrBucketNotFound) { return fmt.Errorf("delete failed for %s: %w", did, err) } } - if err := tx.DeleteBucket([]byte(metaBucket())); err != nil && !errors.Is(err, bbolt.ErrBucketNotFound) { + if err := tx.DeleteBucket([]byte(metaBucket())); err != nil && !errors.Is(err, berr.ErrBucketNotFound) { return fmt.Errorf("delete meta: %w", err) } - if err := tx.DeleteBucket([]byte("quota")); err != nil && !errors.Is(err, bbolt.ErrBucketNotFound) { + if err := tx.DeleteBucket([]byte("quota")); err != nil && !errors.Is(err, berr.ErrBucketNotFound) { return fmt.Errorf("delete quota: %w", err) } @@ -317,17 +318,77 @@ func (s *BoltStorage) Get(key string) (int, error) { return val, err } -func (s *BoltStorage) Set(key string, val int) error { - return s.db.Update(func(tx *bbolt.Tx) error { +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, err := json.Marshal(val) + 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), v) + 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 { + b := tx.Bucket([]byte("quota")) + if b == nil { + return nil + } + for _, key := range keys { + v := b.Get([]byte(key)) + if v == nil { + res[key] = 0 + continue + } + var val int + if err := json.Unmarshal(v, &val); err != nil { + return err + } + res[key] = val + } + return nil + }) + return res, err +} + +func (s *BoltStorage) IncrByMulti(deltas map[string]int) error { + return s.db.Update(func(tx *bbolt.Tx) error { + b, err := tx.CreateBucketIfNotExists([]byte("quota")) + if err != nil { + return err + } + for key, n := range deltas { + var val int + 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 + } + if err := b.Put([]byte(key), newV); err != nil { + return err + } + } + return nil }) } diff --git a/cache/cache_test.go b/cache/cache_test.go index 2222f8d..b85c143 100644 --- a/cache/cache_test.go +++ b/cache/cache_test.go @@ -82,20 +82,3 @@ func TestClear(t *testing.T) { t.Error("cache should be invalid") } } - -func TestQuotaKV(t *testing.T) { - storage := newTestStorage(t) - - err := storage.Set("testkey", 123) - if err != nil { - t.Fatalf("Set failed: %v", err) - } - - val, err := storage.Get("testkey") - if err != nil { - t.Fatalf("Get failed: %v", err) - } - if val != 123 { - t.Errorf("expected 123, got %d", val) - } -} diff --git a/cache/storage.go b/cache/storage.go index 622f495..50c2a0a 100644 --- a/cache/storage.go +++ b/cache/storage.go @@ -26,7 +26,9 @@ type Storage interface { // KVStore implementation Get(key string) (int, error) - Set(key string, val 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.nix b/flake.nix index c79091a..4eefad6 100644 --- a/flake.nix +++ b/flake.nix @@ -17,9 +17,9 @@ let lazuli = pkgs.buildGoModule rec { name = "lazuli"; - version = "0.1.1"; + version = "0.1.2"; src = pkgs.nix-gitignore.gitignoreSource [ "*.csv" "*.zip" "*.json" ] ./.; - vendorHash = "sha256-Zr9gGytJARMbf/7120HYkKsfzpeW47MkwdMODD9QTKc="; + vendorHash = "sha256-MfBPv/L7wHuUGXx4BDd+DFq0RB11KuMHCzPjFv6FMgs="; ldflags = [ "-X" "main.Version=${version}" diff --git a/go.mod b/go.mod index 064937e..c064cb3 100644 --- a/go.mod +++ b/go.mod @@ -3,27 +3,30 @@ module tangled.org/karitham.dev/lazuli go 1.25.5 require ( - github.com/bluesky-social/indigo v0.0.0-20260114211028-207c9d49d0de - github.com/urfave/cli/v3 v3.6.1 - go.etcd.io/bbolt v1.3.10 - golang.org/x/text v0.14.0 - golang.org/x/time v0.3.0 + github.com/bluesky-social/indigo v0.0.0-20260120225912-12d69fa4d209 + github.com/failsafe-go/failsafe-go v0.9.5 + github.com/urfave/cli/v3 v3.6.2 + go.etcd.io/bbolt v1.4.3 + golang.org/x/text v0.33.0 ) require ( github.com/beorn7/perks v1.0.1 // indirect - github.com/cespare/xxhash/v2 v2.2.0 // indirect + github.com/bits-and-blooms/bitset v1.24.4 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/earthboundkid/versioninfo/v2 v2.24.1 // indirect github.com/hashicorp/golang-lru/v2 v2.0.7 // indirect - github.com/matttproud/golang_protobuf_extensions/v2 v2.0.0 // indirect github.com/mr-tron/base58 v1.2.0 // indirect - github.com/prometheus/client_golang v1.17.0 // indirect - github.com/prometheus/client_model v0.5.0 // indirect - github.com/prometheus/common v0.45.0 // indirect - github.com/prometheus/procfs v0.12.0 // indirect + github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect + github.com/prometheus/client_golang v1.23.2 // indirect + github.com/prometheus/client_model v0.6.2 // indirect + github.com/prometheus/common v0.67.5 // indirect + github.com/prometheus/procfs v0.19.2 // indirect gitlab.com/yawning/secp256k1-voi v0.0.0-20230925100816-f2616030848b // indirect gitlab.com/yawning/tuplehash v0.0.0-20230713102510-df83abbf9a02 // indirect - golang.org/x/crypto v0.21.0 // indirect - golang.org/x/sys v0.22.0 // indirect - google.golang.org/protobuf v1.33.0 // indirect + go.yaml.in/yaml/v2 v2.4.3 // indirect + golang.org/x/crypto v0.47.0 // indirect + golang.org/x/sys v0.40.0 // indirect + golang.org/x/time v0.14.0 // indirect + google.golang.org/protobuf v1.36.11 // indirect ) diff --git a/go.sum b/go.sum index 4f17b6d..3662d09 100644 --- a/go.sum +++ b/go.sum @@ -1,23 +1,31 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= -github.com/bluesky-social/indigo v0.0.0-20260114211028-207c9d49d0de h1:75emVEzhTQWXwAQoBZV4/Bg2NEULZSgRwLFAdTccTrY= -github.com/bluesky-social/indigo v0.0.0-20260114211028-207c9d49d0de/go.mod h1:KIy0FgNQacp4uv2Z7xhNkV3qZiUSGuRky97s7Pa4v+o= -github.com/cespare/xxhash/v2 v2.2.0 h1:DC2CZ1Ep5Y4k3ZQ899DldepgrayRUGE6BBZ/cd9Cj44= -github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/bits-and-blooms/bitset v1.24.4 h1:95H15Og1clikBrKr/DuzMXkQzECs1M6hhoGXLwLQOZE= +github.com/bits-and-blooms/bitset v1.24.4/go.mod h1:7hO7Gc7Pp1vODcmWvKMRA9BNmbv6a/7QIWpPxHddWR8= +github.com/bluesky-social/indigo v0.0.0-20260120225912-12d69fa4d209 h1:W01PGqjCexVBzIZ4FoNe4iO8OhI9XbSE7ieWL0QnMu8= +github.com/bluesky-social/indigo v0.0.0-20260120225912-12d69fa4d209/go.mod h1:KIy0FgNQacp4uv2Z7xhNkV3qZiUSGuRky97s7Pa4v+o= +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= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/earthboundkid/versioninfo/v2 v2.24.1 h1:SJTMHaoUx3GzjjnUO1QzP3ZXK6Ee/nbWyCm58eY3oUg= github.com/earthboundkid/versioninfo/v2 v2.24.1/go.mod h1:VcWEooDEuyUJnMfbdTh0uFN4cfEIg+kHMuWB2CDCLjw= -github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= -github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/failsafe-go/failsafe-go v0.9.5 h1:Bgt4wTKV3+n49GssB2njPZ4u5ApjvtKSIQlqIL4E3oo= +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/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= +github.com/influxdata/tdigest v0.0.1/go.mod h1:Z0kXnxzbTC2qrx4NaIzYkE1k66+6oEDQTvL95hQFh5Y= github.com/ipfs/go-cid v0.4.1 h1:A/T3qGvxi4kpKWWcPC/PgbvDA2bjVLO7n4UeVwnbs/s= github.com/ipfs/go-cid v0.4.1/go.mod h1:uQHwDeX4c6CtyrFwdqyhpNcxVewur1M7l7fNU7LKwZk= github.com/klauspost/cpuid/v2 v2.2.7 h1:ZWSB3igEs+d0qvnxR/ZBzXVmxkgt8DdzP6m9pfuVLDM= github.com/klauspost/cpuid/v2 v2.2.7/go.mod h1:Lcz8mBdAVJIBVzewtcLocK12l3Y+JytZYpaMropDUws= -github.com/matttproud/golang_protobuf_extensions/v2 v2.0.0 h1:jWpvCLoY8Z/e3VKvlsiIGKtc+UG6U5vzxaoagmhXfyg= -github.com/matttproud/golang_protobuf_extensions/v2 v2.0.0/go.mod h1:QUyp042oQthUoa9bqDv0ER0wrtXnBruoNd7aNjkbP+k= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/minio/sha256-simd v1.0.1 h1:6kaan5IFmwTNynnKKpDHe6FWHohJOHhCPchzK49dzMM= github.com/minio/sha256-simd v1.0.1/go.mod h1:Pz6AKMiUdngCLpeTL/RJY1M9rUuPMYujV5xJjtbRSN8= github.com/mr-tron/base58 v1.2.0 h1:T/HDJBh4ZCPbU39/+c3rRvE0uKBQlU27+QI8LJ4t64o= @@ -32,44 +40,61 @@ github.com/multiformats/go-multihash v0.2.3 h1:7Lyc8XfX/IY2jWb/gI7JP+o7JEq9hOa7B github.com/multiformats/go-multihash v0.2.3/go.mod h1:dXgKXCXjBzdscBLk9JkjINiEsCKRVch90MdaGiKsvSM= github.com/multiformats/go-varint v0.0.7 h1:sWSGR+f/eu5ABZA2ZpYKBILXTTs9JWpdEM/nEGOHFS8= github.com/multiformats/go-varint v0.0.7/go.mod h1:r8PUYw/fD/SjBCiKOoDlGF6QawOELpZAu9eioSos/OU= +github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= +github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/prometheus/client_golang v1.17.0 h1:rl2sfwZMtSthVU752MqfjQozy7blglC+1SOtjMAMh+Q= -github.com/prometheus/client_golang v1.17.0/go.mod h1:VeL+gMmOAxkS2IqfCq0ZmHSL+LjWfWDUmp1mBz9JgUY= -github.com/prometheus/client_model v0.5.0 h1:VQw1hfvPvk3Uv6Qf29VrPF32JB6rtbgI6cYPYQjL0Qw= -github.com/prometheus/client_model v0.5.0/go.mod h1:dTiFglRmd66nLR9Pv9f0mZi7B7fk5Pm3gvsjB5tr+kI= -github.com/prometheus/common v0.45.0 h1:2BGz0eBc2hdMDLnO/8n0jeB3oPrt2D08CekT0lneoxM= -github.com/prometheus/common v0.45.0/go.mod h1:YJmSTw9BoKxJplESWWxlbyttQR4uaEcGyv9MZjVOJsY= -github.com/prometheus/procfs v0.12.0 h1:jluTpSng7V9hY0O2R9DzzJHYb2xULk9VTR1V1R/k6Bo= -github.com/prometheus/procfs v0.12.0/go.mod h1:pcuDEFsWDnvcgNzo4EEweacyhjeA9Zk3cnaOZAZEfOo= +github.com/prometheus/client_golang v1.23.2 h1:Je96obch5RDVy3FDMndoUsjAhG5Edi49h0RJWRi/o0o= +github.com/prometheus/client_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg= +github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk= +github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE= +github.com/prometheus/common v0.67.5 h1:pIgK94WWlQt1WLwAC5j2ynLaBRDiinoAb86HZHTUGI4= +github.com/prometheus/common v0.67.5/go.mod h1:SjE/0MzDEEAyrdr5Gqc6G+sXI67maCxzaT3A2+HqjUw= +github.com/prometheus/procfs v0.19.2 h1:zUMhqEW66Ex7OXIiDkll3tl9a1ZdilUOd/F6ZXw4Vws= +github.com/prometheus/procfs v0.19.2/go.mod h1:M0aotyiemPhBCM0z5w87kL22CxfcH05ZpYlu+b4J7mw= +github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= +github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= github.com/spaolacci/murmur3 v1.1.0 h1:7c1g84S4BPRrfL5Xrdp6fOJ206sU9y293DDHaoy0bLI= github.com/spaolacci/murmur3 v1.1.0/go.mod h1:JwIasOWyU6f++ZhiEuf87xNszmSA2myDM2Kzu9HwQUA= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= -github.com/urfave/cli/v3 v3.6.1 h1:j8Qq8NyUawj/7rTYdBGrxcH7A/j7/G8Q5LhWEW4G3Mo= -github.com/urfave/cli/v3 v3.6.1/go.mod h1:ysVLtOEmg2tOy6PknnYVhDoouyC/6N42TMeoMzskhso= +github.com/urfave/cli/v3 v3.6.2 h1:lQuqiPrZ1cIz8hz+HcrG0TNZFxU70dPZ3Yl+pSrH9A8= +github.com/urfave/cli/v3 v3.6.2/go.mod h1:ysVLtOEmg2tOy6PknnYVhDoouyC/6N42TMeoMzskhso= github.com/whyrusleeping/cbor-gen v0.2.1-0.20241030202151-b7a6831be65e h1:28X54ciEwwUxyHn9yrZfl5ojgF4CBNLWX7LR0rvBkf4= github.com/whyrusleeping/cbor-gen v0.2.1-0.20241030202151-b7a6831be65e/go.mod h1:pM99HXyEbSQHcosHc0iW7YFmwnscr+t9Te4ibko05so= gitlab.com/yawning/secp256k1-voi v0.0.0-20230925100816-f2616030848b h1:CzigHMRySiX3drau9C6Q5CAbNIApmLdat5jPMqChvDA= gitlab.com/yawning/secp256k1-voi v0.0.0-20230925100816-f2616030848b/go.mod h1:/y/V339mxv2sZmYYR64O07VuCpdNZqCTwO8ZcouTMI8= gitlab.com/yawning/tuplehash v0.0.0-20230713102510-df83abbf9a02 h1:qwDnMxjkyLmAFgcfgTnfJrmYKWhHnci3GjDqcZp1M3Q= gitlab.com/yawning/tuplehash v0.0.0-20230713102510-df83abbf9a02/go.mod h1:JTnUj0mpYiAsuZLmKjTx/ex3AtMowcCgnE7YNyCEP0I= -go.etcd.io/bbolt v1.3.10 h1:+BqfJTcCzTItrop8mq/lbzL8wSGtj94UO/3U31shqG0= -go.etcd.io/bbolt v1.3.10/go.mod h1:bK3UQLPJZly7IlNmV7uVHJDxfe5aK9Ll93e/74Y9oEQ= -golang.org/x/crypto v0.21.0 h1:X31++rzVUdKhX5sWmSOFZxx8UW/ldWx55cbf08iNAMA= -golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOMs= -golang.org/x/sync v0.7.0 h1:YsImfSBoP9QPYL0xyKJPq0gcaJdG3rInoqxTWbfQu9M= -golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= -golang.org/x/sys v0.22.0 h1:RI27ohtqKCnwULzJLqkv897zojh5/DwS/ENaMzUOaWI= -golang.org/x/sys v0.22.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= -golang.org/x/text v0.14.0 h1:ScX5w1eTa3QqT8oi6+ziP7dTV1S2+ALU0bI+0zXKWiQ= -golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= -golang.org/x/time v0.3.0 h1:rg5rLMjNzMS1RkNLzCG38eapWhnYLFYXDXj2gOlr8j4= -golang.org/x/time v0.3.0/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ= +go.etcd.io/bbolt v1.4.3 h1:dEadXpI6G79deX5prL3QRNP6JB8UxVkqo4UPnHaNXJo= +go.etcd.io/bbolt v1.4.3/go.mod h1:tKQlpPaYCVFctUIgFKFnAlvbmB3tpy1vkTnDWohtc0E= +go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= +go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= +go.yaml.in/yaml/v2 v2.4.3 h1:6gvOSjQoTB3vt1l+CU+tSyi/HOjfOjRLJ4YwYZGwRO0= +go.yaml.in/yaml/v2 v2.4.3/go.mod h1:zSxWcmIDjOzPXpjlTTbAsKokqkDNAVtZO0WOMiT90s8= +golang.org/x/crypto v0.47.0 h1:V6e3FRj+n4dbpw86FJ8Fv7XVOql7TEwpHapKoMJ/GO8= +golang.org/x/crypto v0.47.0/go.mod h1:ff3Y9VzzKbwSSEzWqJsJVBnWmRwRSHt/6Op5n9bQc4A= +golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU= +golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY= +golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= +golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= +golang.org/x/sys v0.40.0 h1:DBZZqJ2Rkml6QMQsZywtnjnnGvHza6BTfYFWY9kjEWQ= +golang.org/x/sys v0.40.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/text v0.33.0 h1:B3njUFyqtHDUI5jMn1YIr5B0IE2U0qck04r6d4KPAxE= +golang.org/x/text v0.33.0/go.mod h1:LuMebE6+rBincTi9+xWTY8TztLzKHc/9C1uBCG27+q8= +golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI= +golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4= golang.org/x/xerrors v0.0.0-20231012003039-104605ab7028 h1:+cNy6SZtPcJQH3LJVLOSmiC7MMxXNOb3PU/VUEz+EhU= golang.org/x/xerrors v0.0.0-20231012003039-104605ab7028/go.mod h1:NDW/Ps6MPRej6fsCIbMTohpP40sJ/P/vI1MoTEGwX90= -google.golang.org/protobuf v1.33.0 h1:uNO2rsAINq/JlFpSdYEKIZ0uKD/R9cpdv0T+yoGwGmI= -google.golang.org/protobuf v1.33.0/go.mod h1:c6P6GXX6sHbq/GpV6MGZEdwhWPcYBgnhAHhKbcUYpos= +google.golang.org/genproto/googleapis/rpc v0.0.0-20240814211410-ddb44dafa142 h1:e7S5W7MGGLaSu8j3YjdezkZ+m1/Nm0uRVRMEMGk26Xs= +google.golang.org/genproto/googleapis/rpc v0.0.0-20240814211410-ddb44dafa142/go.mod h1:UqMtugtsSgubUsoxbuAoiCXvqvErP7Gf0so0mK9tHxU= +google.golang.org/grpc v1.67.1 h1:zWnc1Vrcno+lHZCOofnIMvycFcc0QRGIzm9dhnDX68E= +google.golang.org/grpc v1.67.1/go.mod h1:1gLDyUQU7CTLJI90u3nXZ9ekeghjeM7pTDZlqFNg2AA= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= lukechampine.com/blake3 v1.2.1 h1:YuqqRuaqsGV71BV/nm9xlI0MKUv4QC54jQnBChWbGnI= diff --git a/main.go b/main.go index a74da1d..66b25b7 100644 --- a/main.go +++ b/main.go @@ -17,6 +17,8 @@ import ( "tangled.org/karitham.dev/lazuli/sync" "tangled.org/karitham.dev/lazuli/sync/logutil" + "github.com/failsafe-go/failsafe-go" + "github.com/failsafe-go/failsafe-go/retrypolicy" "github.com/urfave/cli/v3" ) @@ -159,7 +161,10 @@ func (a *App) runStats(ctx context.Context, cmd *cli.Command) error { } limiter := sync.NewRateLimiter(a.storage) - writes, global := limiter.Stats() + writes, global, err := limiter.Stats() + if err != nil { + return fmt.Errorf("failed to get rate limit stats: %w", err) + } if a.outputFormat == "json" { out := map[string]any{ @@ -283,14 +288,10 @@ func (a *App) runRetry(ctx context.Context, cmd *cli.Command) error { } // Check rate limit for 1 write - if err := limiter.AllowBulkWrite(ctx, 1); err != nil { - return fmt.Errorf("rate limit wait failed: %w", err) - } - w, g := limiter.Stats() - res := sync.PublishBatch(ctx, repoClient, did, []sync.PlayRecord{fr.rec}, w, g, a.storage) + res := sync.PublishBatch(ctx, repoClient, did, []sync.PlayRecord{fr.rec}, a.storage) - if res.ErrorCount == 0 { + if res == nil { fmt.Printf("Successfully retried: %s - %s\n", fr.rec.ArtistName(), fr.rec.TrackName) // Mark as published (updates processedBucket to 1) if err := a.storage.MarkPublished(did, fr.key); err != nil { @@ -302,9 +303,8 @@ func (a *App) runRetry(ctx context.Context, cmd *cli.Command) error { } successCount++ } else { - fmt.Printf("Failed again: %s - %s: %v\n", fr.rec.ArtistName(), fr.rec.TrackName, res.LastError) + fmt.Printf("Failed again: %s - %s: %v\n", fr.rec.ArtistName(), fr.rec.TrackName, res) errorCount++ - limiter.RefundBulkWrite(1) } // Optional: small delay between retries? @@ -600,9 +600,10 @@ func (a *App) runImport(ctx context.Context, cmd *cli.Command) error { ATProtoClient: repoClient, ProgressLog: progressLog, Storage: a.storage, + Limiter: limiter, } - result := sync.Publish(ctx, authClient, publishOpts, limiter) + result := sync.Publish(ctx, authClient, publishOpts) a.log.Info("Import completed", slog.Int("success_count", result.SuccessCount), @@ -646,7 +647,9 @@ func (a *App) createProgressLogger() func(sync.ProgressReport) { slog.String("elapsed", pr.Elapsed), slog.String("eta", pr.ETA), slog.String("rate", pr.Rate), - slog.Int("errors", pr.Errors)) + slog.Int("errors", pr.Errors), + slog.Int("writes", pr.WritesConsumed), + slog.Int("global", pr.GlobalConsumed)) } } } @@ -758,7 +761,26 @@ func (a *App) runDedupe(ctx context.Context, cmd *cli.Command) error { uri := rec.URI parts := strings.Split(uri, "/") rkey := parts[len(parts)-1] - err := repoClient.DeleteRecord(ctx, sync.RecordType, rkey) + + retryPolicy := retrypolicy.NewBuilder[any](). + WithMaxRetries(10). + WithBackoff(sync.BaseRetryDelay, 5*time.Minute). + HandleIf(func(_ any, err error) bool { + return sync.IsTransientError(err) + }). + OnRetryScheduled(func(e failsafe.ExecutionScheduledEvent[any]) { + a.log.Warn("Delete failed with transient error, retrying", + slog.Duration("retryDelay", e.Delay), + logutil.Error(e.LastError()), + slog.Int("attempt", e.Attempts()), + slog.String("uri", uri)) + }). + Build() + + err := failsafe.With[any](retryPolicy).WithContext(ctx).Run(func() error { + return repoClient.DeleteRecord(ctx, sync.RecordType, rkey) + }) + if err != nil { a.log.Error("Failed to delete record", logutil.Error(err), slog.String("uri", uri)) } else { diff --git a/sync/adapter.go b/sync/adapter.go index 12c34dc..693c687 100644 --- a/sync/adapter.go +++ b/sync/adapter.go @@ -3,12 +3,18 @@ package sync import ( "context" "encoding/json" + "errors" "fmt" "log/slog" + "math/rand" + "net" "strings" + "time" "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/syntax" + "github.com/failsafe-go/failsafe-go" + "github.com/failsafe-go/failsafe-go/retrypolicy" "tangled.org/karitham.dev/lazuli/cache" "tangled.org/karitham.dev/lazuli/sync/logutil" @@ -45,70 +51,62 @@ func (c *RateClient) ListRecords(ctx context.Context, collection string, limit i return nil, "", fmt.Errorf("client cannot be nil") } + var outResp struct { + Records []struct { + URI string `json:"uri"` + CID string `json:"cid"` + Value map[string]any `json:"value"` + } `json:"records"` + Cursor string `json:"cursor"` + } + + var chargedAt time.Time if c.limiter != nil { slog.Debug("waiting for rate limit (read)") - if err := c.limiter.AllowRead(ctx); err != nil { + var err error + chargedAt, err = c.limiter.AllowRead(ctx) + if err != nil { slog.Error("rate limit wait cancelled/failed (read)", logutil.Error(err)) return nil, "", err } } - var out []RecordRef - - for { - select { - case <-ctx.Done(): - return nil, "", ctx.Err() - default: - } - - var outResp struct { - Records []struct { - URI string `json:"uri"` - CID string `json:"cid"` - Value map[string]any `json:"value"` - } `json:"records"` - Cursor string `json:"cursor"` - } - - err := c.client.Get(ctx, syntax.NSID("com.atproto.repo.listRecords"), map[string]any{ - "repo": c.did, - "collection": collection, - "limit": limit, - "cursor": cursor, - }, &outResp) - if err != nil { - return nil, "", err + err := c.client.Get(ctx, syntax.NSID("com.atproto.repo.listRecords"), map[string]any{ + "repo": c.did, + "collection": collection, + "limit": limit, + "cursor": cursor, + }, &outResp) + if err != nil { + if c.limiter != nil && isTransientError(err) { + c.limiter.RefundRead(ctx, chargedAt) } + return nil, "", err + } - for _, r := range outResp.Records { - var playRecord PlayRecord - if r.Value != nil { - b, err := json.Marshal(r.Value) - if err != nil { - slog.Debug("failed to marshal record value", slog.String("uri", r.URI), logutil.Error(err)) - continue - } - if err := json.Unmarshal(b, &playRecord); err != nil { - slog.Debug("failed to unmarshal record", slog.String("uri", r.URI), logutil.Error(err)) - continue - } - slog.Debug("parsed record", slog.String("uri", r.URI), logutil.Track(playRecord.TrackName, playRecord.ArtistName(), playRecord.PlayedTime.Time)) + out := make([]RecordRef, 0, len(outResp.Records)) + for _, r := range outResp.Records { + var playRecord PlayRecord + if r.Value != nil { + b, err := json.Marshal(r.Value) + if err != nil { + slog.Debug("failed to marshal record value", slog.String("uri", r.URI), logutil.Error(err)) + continue } - out = append(out, RecordRef{ - URI: r.URI, - CID: r.CID, - Value: playRecord, - }) - } - - if outResp.Cursor == "" || len(outResp.Records) < limit { - break + if err := json.Unmarshal(b, &playRecord); err != nil { + slog.Debug("failed to unmarshal record", slog.String("uri", r.URI), logutil.Error(err)) + continue + } + slog.Debug("parsed record", slog.String("uri", r.URI), logutil.Track(playRecord.TrackName, playRecord.ArtistName(), playRecord.PlayedTime.Time)) } - cursor = outResp.Cursor + out = append(out, RecordRef{ + URI: r.URI, + CID: r.CID, + Value: playRecord, + }) } - return out, cursor, nil + return out, outResp.Cursor, nil } func (c *RateClient) ApplyWrites(ctx context.Context, collection string, records []PlayRecord) error { @@ -120,27 +118,248 @@ func (c *RateClient) ApplyWrites(ctx context.Context, collection string, records return fmt.Errorf("client cannot be nil") } + var chargedAt time.Time if c.limiter != nil { - slog.Debug("waiting for rate limit (write)") - if err := c.limiter.AllowBulkWrite(ctx, len(records)); err != nil { - slog.Error("rate limit wait cancelled/failed (write)", logutil.Error(err)) + var err error + chargedAt, err = c.limiter.AllowBulkWrite(ctx, len(records)) + if err != nil { + slog.Error("rate limit wait cancelled/failed (write)", + logutil.DID(c.did), + slog.String("collection", collection), + logutil.Error(err), + ) return err } } + err := applyWrites(ctx, c.client, c.did, collection, records) + if err != nil && isTransientError(err) && c.limiter != nil { + c.limiter.RefundBulkWrite(ctx, len(records), chargedAt) + } + return err +} + +func applyWrites(ctx context.Context, client *atclient.APIClient, did, collection string, records []PlayRecord) error { + if len(records) == 0 { + return nil + } + + if len(records) > 200 { + return fmt.Errorf("too many records in one ApplyWrites call: %d (max 200)", len(records)) + } + writes, err := prepareWrites(records, collection) if err != nil { return err } - err = c.client.Post(ctx, syntax.NSID("com.atproto.repo.applyWrites"), map[string]any{ - "repo": c.did, + return client.Post(ctx, syntax.NSID("com.atproto.repo.applyWrites"), map[string]any{ + "repo": did, "writes": writes, }, nil) - if err != nil && c.limiter != nil { - c.limiter.RefundBulkWrite(len(records)) +} + +type atprotoClientAdapter struct { + client *atclient.APIClient + did string +} + +func (a *atprotoClientAdapter) ApplyWrites(ctx context.Context, collection string, records []PlayRecord) error { + return applyWrites(ctx, a.client, a.did, collection, records) +} + +func prepareRecords(batch []PlayRecord) []PlayRecord { + atprotoRecords := make([]PlayRecord, 0, len(batch)) + for _, record := range batch { + record.Type = RecordType + record.SubmissionClientAgent = ClientAgent + atprotoRecords = append(atprotoRecords, record) } - return err + return atprotoRecords +} + +func waitForRetry(ctx context.Context, delay time.Duration) bool { + timer := time.NewTimer(delay) + defer timer.Stop() + + select { + case <-timer.C: + return true + case <-ctx.Done(): + slog.Debug("retry cancelled due to context done") + return false + } +} + +func defaultProgressLog(f func(ProgressReport)) func(ProgressReport) { + if f != nil { + return f + } + return func(pr ProgressReport) { + slog.Info("sync progress", + slog.Int("completed", pr.Completed), + slog.Int("total", pr.Total), + slog.Float64("percent", pr.Percent), + slog.String("elapsed", pr.Elapsed), + slog.String("eta", pr.ETA), + slog.String("rate", pr.Rate), + slog.Int("errors", pr.Errors), + ) + } +} + +func defaultBatchSize(size int) int { + if size > 0 { + return size + } + return DefaultBatchSize +} + +func buildClient(client AuthClient, customClient ATProtoClient) (ATProtoClient, error) { + if customClient != nil { + return customClient, nil + } + + apiClient := client.GetAPIClient() + if apiClient == nil { + slog.Error("failed to get API client", logutil.Error(fmt.Errorf("client is nil"))) + return nil, fmt.Errorf("API client is nil") + } + + return &atprotoClientAdapter{client: apiClient, did: client.GetDID()}, nil +} + +func newPublishResult(success, errors, total int, start time.Time, cancelled bool) PublishResult { + return PublishResult{ + SuccessCount: success, + ErrorCount: errors, + Cancelled: cancelled, + Duration: time.Since(start), + TotalRecords: total, + RecordsPerMinute: ratePerMinute(success, time.Since(start)), + } +} + +func logResult(success, errors int, startTime time.Time) { + if errors > 0 { + slog.Warn("import completed with errors", + slog.Int("success", success), + slog.Int("errors", errors)) + } + slog.Info("import completed", + slog.Int("success", success), + slog.Int("errors", errors), + slog.Duration("duration", time.Since(startTime)), + slog.String("rate", formatRate(ratePerMinute(success, time.Since(startTime))))) +} + +func backoff(attempt int) time.Duration { + if attempt <= 0 { + return BaseRetryDelay + } + + // Calculate exponential delay: BaseRetryDelay * 2^(attempt-1) + // We use uint(attempt-1) because 1<<0 is 1 (for first retry) + exp := min(attempt-1, 31) + + delay := BaseRetryDelay * time.Duration(1< MaxRetryDelay || delay <= 0 { + delay = MaxRetryDelay + } + + // Add up to 25% jitter + var jitter time.Duration + if delay > 4 { + jitter = time.Duration(rand.Int63n(int64(delay / 4))) + } + + return delay + jitter +} + +func PublishBatch(ctx context.Context, client ATProtoClient, did string, batch []PlayRecord, storage cache.Storage) error { + if len(batch) == 0 { + return nil + } + + atprotoRecords := prepareRecords(batch) + + err := client.ApplyWrites(ctx, RecordType, atprotoRecords) + if err != nil { + slog.Error("batch publish failed", logutil.Error(err)) + return err + } + + if storage != nil && did != "" { + keys := CreateRecordKeys(atprotoRecords) + cacheEntries := make(map[string][]byte) + for i, rec := range atprotoRecords { + key := keys[i] + value, _ := json.Marshal(rec) + cacheEntries[key] = value + } + + if err := storage.SaveRecords(did, cacheEntries); err != nil { + return fmt.Errorf("failed to save records to storage: %w", err) + } + + if err := storage.MarkPublished(did, keys...); err != nil { + return fmt.Errorf("failed to mark records as published: %w", err) + } + } + + return nil +} + +type ATProtoClient interface { + ApplyWrites(ctx context.Context, collection string, records []PlayRecord) error +} + +func ratePerMinute(count int, duration time.Duration) float64 { + if duration == 0 { + return 0 + } + return float64(count) / duration.Minutes() +} + +type AuthClient interface { + GetAPIClient() *atclient.APIClient + GetDID() string +} + +func IsTransientError(err error) bool { + return isTransientError(err) +} + +func Backoff(attempt int) time.Duration { + return backoff(attempt) +} + +func WaitForRetry(ctx context.Context, delay time.Duration) bool { + return waitForRetry(ctx, delay) +} + +func isTransientError(err error) bool { + if err == nil { + return false + } + + var apiErr *atclient.APIError + if errors.As(err, &apiErr) { + switch apiErr.StatusCode { + case 429, 500, 502, 503, 504: + return true + } + return false + } + + var netErr net.Error + if errors.As(err, &netErr) { + return netErr.Timeout() + } + + return false } func (c *RateClient) DeleteRecord(ctx context.Context, collection, rkey string) error { @@ -148,9 +367,12 @@ func (c *RateClient) DeleteRecord(ctx context.Context, collection, rkey string) return fmt.Errorf("client is nil") } + var chargedAt time.Time if c.limiter != nil { slog.Debug("waiting for rate limit (delete)") - if err := c.limiter.AllowBulkWrite(ctx, 1); err != nil { + var err error + chargedAt, err = c.limiter.AllowBulkWrite(ctx, 1) + if err != nil { slog.Error("rate limit wait cancelled/failed (delete)", logutil.Error(err)) return err } @@ -165,8 +387,8 @@ func (c *RateClient) DeleteRecord(ctx context.Context, collection, rkey string) "rkey": {rkey}, }, }) - if err != nil && c.limiter != nil { - c.limiter.RefundBulkWrite(1) + if err != nil && c.limiter != nil && isTransientError(err) { + c.limiter.RefundBulkWrite(ctx, 1, chargedAt) } return err } @@ -201,9 +423,32 @@ func FetchExisting(ctx context.Context, client RepoClient, did string, storage c } allRecords := make([]ExistingRecord, 0, 1024) + return fetchExistingLoop(ctx, client, did, storage, allRecords) +} + +func fetchExistingLoop(ctx context.Context, client RepoClient, did string, storage cache.Storage, allRecords []ExistingRecord) ([]ExistingRecord, error) { const batchSize = 100 var cursor string + type fetchResult struct { + records []RecordRef + cursor string + } + + retryPolicy := retrypolicy.NewBuilder[fetchResult](). + WithMaxRetries(10). + WithBackoff(BaseRetryDelay, 5*time.Minute). + HandleIf(func(_ fetchResult, err error) bool { + return isTransientError(err) + }). + OnRetryScheduled(func(e failsafe.ExecutionScheduledEvent[fetchResult]) { + slog.Warn("fetch failed with transient error, retrying", + slog.Duration("retryDelay", e.Delay), + logutil.Error(e.LastError()), + slog.Int("attempt", e.Attempts())) + }). + Build() + for { select { case <-ctx.Done(): @@ -211,27 +456,34 @@ func FetchExisting(ctx context.Context, client RepoClient, did string, storage c default: } - records, newCursor, err := client.ListRecords(ctx, RecordType, batchSize, cursor) + result, err := failsafe.With(retryPolicy). + WithContext(ctx). + Get(func() (fetchResult, error) { + recs, next, err := client.ListRecords(ctx, RecordType, batchSize, cursor) + if err != nil { + return fetchResult{}, err + } + + return fetchResult{records: recs, cursor: next}, nil + }) if err != nil { return nil, err } - for _, rec := range records { + for _, rec := range result.records { allRecords = append(allRecords, ExistingRecord(rec)) } - if newCursor == "" || len(records) < batchSize { + if result.cursor == "" || len(result.records) < batchSize { break } - cursor = newCursor + cursor = result.cursor } if storage != nil { cacheEntries := make(map[string][]byte) keys := make([]string, 0, len(allRecords)) for _, rec := range allRecords { - // use the rkey from URI if available, otherwise fallback to generating it - // URI is at://did/collection/rkey parts := strings.Split(rec.URI, "/") key := parts[len(parts)-1] if key == "" { @@ -246,7 +498,6 @@ func FetchExisting(ctx context.Context, client RepoClient, did string, storage c return nil, err } - // Mark remote records as published locally to prevent redundant syncs if err := storage.MarkPublished(did, keys...); err != nil { return nil, err } diff --git a/sync/batch_test.go b/sync/batch_test.go new file mode 100644 index 0000000..1be2b58 --- /dev/null +++ b/sync/batch_test.go @@ -0,0 +1,364 @@ +package sync + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/bluesky-social/indigo/atproto/atclient" + "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() + for k, v := range records { + m.unpublished[k] = v + } + 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 +} + +// Mock ATProtoClient +type mockATProtoClient struct { + applyWritesFunc func(ctx context.Context, collection string, records []PlayRecord) 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 +} + +// Mock AuthClient +type mockAuthClient struct { + did string +} + +func (m *mockAuthClient) GetAPIClient() *atclient.APIClient { return nil } +func (m *mockAuthClient) GetDID() 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) { + 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() + + t.Run("Empty records", func(t *testing.T) { + limiter := &mockLimiter{} + client := NewRateClient(nil, "did:example:123", limiter) + err := client.ApplyWrites(ctx, "test", nil) + if err != nil { + t.Errorf("ApplyWrites(nil) error = %v", err) + } + }) + + t.Run("Too many records", func(t *testing.T) { + limiter := &mockLimiter{} + client := NewRateClient(&atclient.APIClient{}, "did:example:123", limiter) + records := make([]PlayRecord, 201) + err := client.ApplyWrites(ctx, "test", records) + if err == nil { + t.Fatal("expected error for > 200 records") + } + expected := "too many records in one ApplyWrites call: 201 (max 200)" + if err.Error() != expected { + t.Errorf("expected error %q, got %q", expected, err.Error()) + } + }) + + t.Run("Transient error refunds tokens", func(t *testing.T) { + 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 + }) + client := NewRateClient(apiClient, "did:example:123", limiter) + + err := client.ApplyWrites(ctx, "test", []PlayRecord{{TrackName: "Song 1"}}) + if err == nil { + t.Fatal("expected error") + } + if atomic.LoadInt32(&limiter.refunds) != 1 { + t.Errorf("expected 1 refund, got %d", limiter.refunds) + } + }) + + t.Run("Non-transient error does NOT refund", func(t *testing.T) { + 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 + }) + client := NewRateClient(apiClient, "did:example:123", limiter) + + err := client.ApplyWrites(ctx, "test", []PlayRecord{{TrackName: "Song 1"}}) + if err == nil { + t.Fatal("expected error") + } + if atomic.LoadInt32(&limiter.refunds) != 0 { + t.Errorf("expected 0 refunds, got %d", limiter.refunds) + } + }) +} + +func TestPublishBatch(t *testing.T) { + ctx := context.Background() + did := "did:example:123" + batch := []PlayRecord{{TrackName: "Song 1"}} + + t.Run("Success", func(t *testing.T) { + storage := newMockStorage() + client := &mockATProtoClient{} + err := PublishBatch(ctx, client, did, batch, storage) + 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) + 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) + 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" + + rec1, _ := json.Marshal(PlayRecord{TrackName: "Song 1"}) + rec2, _ := json.Marshal(PlayRecord{TrackName: "Song 2"}) + + t.Run("Retry on transient error", 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 + }, + } + + oldBase := BaseRetryDelay + BaseRetryDelay = time.Millisecond + defer func() { BaseRetryDelay = oldBase }() + + res := Publish(ctx, &mockAuthClient{did: did}, PublishOptions{ + BatchSize: 1, + ATProtoClient: client, + Storage: storage, + }) + + 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, + }) + + 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/config.go b/sync/config.go index 11d4546..044b90f 100644 --- a/sync/config.go +++ b/sync/config.go @@ -13,10 +13,11 @@ const ( CacheVersion = 1 SlingshotResolverURL = "https://slingshot.microcosm.blue/xrpc/com.bad-example.identity.resolveMiniDoc" MaxRetryDelay = 15 * time.Minute - BaseRetryDelay = 2 * time.Second MaxRetries = 1000 ) +var BaseRetryDelay = 2 * time.Second + type ImportMode string const ( diff --git a/sync/import_test.go b/sync/import_test.go index 9a8e75d..03c63b0 100644 --- a/sync/import_test.go +++ b/sync/import_test.go @@ -3,6 +3,7 @@ package sync_test import ( "context" "encoding/json" + "fmt" "io" "io/fs" "testing" @@ -36,6 +37,9 @@ func (m *mockRepoClient) ListRecords(ctx context.Context, collection string, lim } func (m *mockRepoClient) ApplyWrites(ctx context.Context, collection string, records []sync.PlayRecord) error { + if len(records) > 200 { + return fmt.Errorf("too many records") + } m.applied = append(m.applied, records...) return nil } @@ -52,6 +56,39 @@ type mockAuthClient struct { func (m *mockAuthClient) GetAPIClient() *atclient.APIClient { return nil } func (m *mockAuthClient) GetDID() string { return m.did } +type mockKV struct { + data map[string]int +} + +func (m *mockKV) GetMulti(keys []string) (map[string]int, error) { + out := make(map[string]int) + for _, k := range keys { + out[k] = m.data[k] + } + return out, nil +} + +func (m *mockKV) IncrByMulti(counts map[string]int) error { + for k, v := range counts { + m.data[k] += v + } + return nil +} + +func (m *mockKV) Get(key string) (int, error) { + return m.data[key], nil +} + +func (m *mockKV) Set(key string, val int) error { + m.data[key] = val + return nil +} + +func (m *mockKV) IncrBy(key string, n int) (int, error) { + m.data[key] += n + return m.data[key], nil +} + func TestImportE2E(t *testing.T) { ctx := context.Background() did := "did:plc:test" @@ -124,15 +161,17 @@ func TestImportE2E(t *testing.T) { } // 7. Publish - limiter := sync.NewRateLimiter(storage) + kv := &mockKV{data: make(map[string]int)} + limiter := sync.NewRateLimiter(kv) publishOpts := sync.PublishOptions{ BatchSize: 10, ATProtoClient: mockRepo, Storage: storage, + Limiter: limiter, } auth := &mockAuthClient{did: did} - result := sync.Publish(ctx, auth, publishOpts, limiter) + result := sync.Publish(ctx, auth, publishOpts) if result.SuccessCount != 1 { t.Errorf("expected 1 successful publish, got %d", result.SuccessCount) diff --git a/sync/progress.go b/sync/progress.go index 1e6ae57..362e1d0 100644 --- a/sync/progress.go +++ b/sync/progress.go @@ -118,15 +118,17 @@ type ProgressTracker struct { LastLogTime time.Time mu sync.Mutex + limiter RateLimiter LogInterval time.Duration LogRecordsMetric int } -func NewProgressTracker(total int) *ProgressTracker { +func NewProgressTracker(total int, limiter RateLimiter) *ProgressTracker { return &ProgressTracker{ Total: total, StartTime: time.Now(), LastLogTime: time.Now(), + limiter: limiter, LogInterval: 30 * time.Second, LogRecordsMetric: 1000, } @@ -191,13 +193,15 @@ func (t *ProgressTracker) ShouldLog() bool { } type ProgressReport struct { - Total int `json:"total"` - Completed int `json:"completed"` - Percent float64 `json:"percent"` - Errors int `json:"errors"` - Elapsed string `json:"elapsed"` - ETA string `json:"eta,omitempty"` - Rate string `json:"rate"` + Total int `json:"total"` + Completed int `json:"completed"` + Percent float64 `json:"percent"` + Errors int `json:"errors"` + Elapsed string `json:"elapsed"` + ETA string `json:"eta,omitempty"` + Rate string `json:"rate"` + WritesConsumed int `json:"writesConsumed,omitempty"` + GlobalConsumed int `json:"globalConsumed,omitempty"` } func (t *ProgressTracker) Report() ProgressReport { @@ -206,14 +210,22 @@ func (t *ProgressTracker) Report() ProgressReport { if eta > 0 { etaStr = FormatDuration(eta) } + + var w, g int + if t.limiter != nil { + w, g, _ = t.limiter.Stats() + } + return ProgressReport{ - Total: t.Total, - Completed: t.Completed, - Percent: percent, - Errors: t.Errors, - Elapsed: elapsed.Round(time.Second).String(), - ETA: etaStr, - Rate: rate, + Total: t.Total, + Completed: t.Completed, + Percent: percent, + Errors: t.Errors, + Elapsed: elapsed.Round(time.Second).String(), + ETA: etaStr, + Rate: rate, + WritesConsumed: w, + GlobalConsumed: g, } } diff --git a/sync/publish.go b/sync/publish.go index befed71..7fc0e57 100644 --- a/sync/publish.go +++ b/sync/publish.go @@ -3,14 +3,13 @@ package sync import ( "context" "encoding/json" - "errors" "fmt" "log/slog" - "math/rand" "time" - "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/syntax" + "github.com/failsafe-go/failsafe-go" + "github.com/failsafe-go/failsafe-go/retrypolicy" "tangled.org/karitham.dev/lazuli/cache" "tangled.org/karitham.dev/lazuli/sync/logutil" @@ -22,9 +21,10 @@ type PublishOptions struct { ATProtoClient ATProtoClient ProgressLog func(ProgressReport) Storage cache.Storage + Limiter RateLimiter } -func Publish(ctx context.Context, client AuthClient, opts PublishOptions, limiter RateLimiter) PublishResult { +func Publish(ctx context.Context, client AuthClient, opts PublishOptions) PublishResult { startTime := time.Now() batchSize := defaultBatchSize(opts.BatchSize) @@ -59,7 +59,7 @@ func Publish(ctx context.Context, client AuthClient, opts PublishOptions, limite slog.Int("daily_token_limit", GlobalLimitDay), slog.String("rate_limit", fmt.Sprintf("1 write per %.1fs", 86400.0/WriteLimitDay))) - tracker := NewProgressTracker(totalRecords) + tracker := NewProgressTracker(totalRecords, opts.Limiter) progressLog := defaultProgressLog(opts.ProgressLog) totalSuccess := 0 totalErrors := 0 @@ -86,78 +86,46 @@ func Publish(ctx context.Context, client AuthClient, opts PublishOptions, limite return nil } - if err := limiter.AllowBulkWrite(ctx, len(batch)); err != nil { - slog.Error("rate limit wait failed", logutil.Error(err)) - return err - } - - var lastResult BatchResult did := client.GetDID() - attempt := 0 - for { - w, g := limiter.Stats() - lastResult = PublishBatch(ctx, atprotoClient, did, batch, w, g, opts.Storage) - - if lastResult.ErrorCount == 0 { - break - } - - limiter.RefundBulkWrite(len(batch)) - - attempt++ - var apiErr *atclient.APIError - is500 := errors.As(lastResult.LastError, &apiErr) && apiErr.StatusCode >= 500 - - if is500 || attempt >= MaxRetries { - if is500 { - first := batch[0] - last := batch[len(batch)-1] - slog.Error("batch failed with 500 error, marking as failed and moving on", - logutil.Error(lastResult.LastError), - slog.Int("count", len(batch)), - slog.Group("range", - slog.Attr(logutil.Track(first.TrackName, first.ArtistName(), first.PlayedTime.Time)), - slog.Attr(logutil.Track(last.TrackName, last.ArtistName(), last.PlayedTime.Time)))) - } else { - slog.Error("batch failed after max retries", slog.Int("errorCount", lastResult.ErrorCount)) - } - - if opts.Storage != nil { - errMsg := lastResult.LastError.Error() - if err := opts.Storage.MarkFailed(did, batchKeys, errMsg); err != nil { - slog.Error("failed to mark records as failed", logutil.Error(err)) - } + retryPolicy := retrypolicy.NewBuilder[any](). + WithMaxRetries(10). + WithBackoff(BaseRetryDelay, 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), + logutil.Error(e.LastError()), + slog.Int("attempt", e.Attempts())) + }). + Build() + + err := failsafe.With[any](retryPolicy).WithContext(ctx).Run(func() error { + return PublishBatch(ctx, atprotoClient, did, batch, opts.Storage) + }) + if err != nil { + slog.Error("batch failed after retries", + logutil.Error(err), + slog.Int("count", len(batch))) + + if opts.Storage != nil { + if markErr := opts.Storage.MarkFailed(did, batchKeys, err.Error()); markErr != nil { + slog.Error("failed to mark records as failed", logutil.Error(markErr)) } - break } - delay := backoff(attempt) - slog.Warn("batch failed, retrying with backoff", - slog.Int("errorCount", lastResult.ErrorCount), - slog.Duration("retryDelay", delay), - logutil.Error(lastResult.LastError), - slog.Int("attempt", attempt)) - - if !waitForRetry(ctx, delay) { - return ctx.Err() - } + totalErrors += len(batch) + tracker.IncrementErrors(len(batch)) - if err := limiter.AllowBulkWrite(ctx, len(batch)); err != nil { - return err - } + batch = batch[:0] + batchKeys = batchKeys[:0] + return nil // Return nil so we continue with the next batch } - totalSuccess += lastResult.SuccessCount - totalErrors += lastResult.ErrorCount - tracker.Increment(lastResult.SuccessCount) - tracker.IncrementErrors(lastResult.ErrorCount) - - if lastResult.SuccessCount > 0 && opts.Storage != nil { - keys := CreateRecordKeys(batch[:lastResult.SuccessCount]) - if err := opts.Storage.MarkPublished(did, keys...); err != nil { - slog.Error("failed to mark records as published", logutil.Error(err)) - } - } + totalSuccess += len(batch) + tracker.Increment(len(batch)) if tracker.ShouldLog() { progressLog(tracker.Report()) @@ -177,7 +145,13 @@ func Publish(ctx context.Context, client AuthClient, opts PublishOptions, limite var record PlayRecord if err := json.Unmarshal(rec, &record); err != nil { - return nil // skip malformed + slog.Error("malformed record in storage", slog.String("key", key), logutil.Error(err)) + if opts.Storage != nil { + _ = opts.Storage.MarkFailed(client.GetDID(), []string{key}, "malformed record") + } + totalErrors++ + tracker.IncrementErrors(1) + return nil } batch = append(batch, record) @@ -204,202 +178,3 @@ func Publish(ctx context.Context, client AuthClient, opts PublishOptions, limite logResult(totalSuccess, totalErrors, startTime) return newPublishResult(totalSuccess, totalErrors, totalRecords, startTime, cancelled) } - -// waitForRetry waits for the specified duration, returning false if context is cancelled. -func waitForRetry(ctx context.Context, delay time.Duration) bool { - select { - case <-time.After(delay): - return true - case <-ctx.Done(): - slog.Debug("retry cancelled due to context done") - return false - } -} - -func defaultProgressLog(f func(ProgressReport)) func(ProgressReport) { - if f != nil { - return f - } - return func(pr ProgressReport) { - slog.Info("sync progress", - slog.Int("completed", pr.Completed), - slog.Int("total", pr.Total), - slog.Float64("percent", pr.Percent), - slog.String("elapsed", pr.Elapsed), - slog.String("eta", pr.ETA), - slog.String("rate", pr.Rate), - slog.Int("errors", pr.Errors), - ) - } -} - -func defaultBatchSize(size int) int { - if size > 0 { - return size - } - return DefaultBatchSize -} - -func buildClient(client AuthClient, customClient ATProtoClient) (ATProtoClient, error) { - if customClient != nil { - return customClient, nil - } - - apiClient := client.GetAPIClient() - if apiClient == nil { - slog.Error("failed to get API client", logutil.Error(fmt.Errorf("client is nil"))) - return nil, fmt.Errorf("API client is nil") - } - - return &atprotoClientAdapter{client: apiClient, did: client.GetDID()}, nil -} - -func newPublishResult(success, errors, total int, start time.Time, cancelled bool) PublishResult { - return PublishResult{ - SuccessCount: success, - ErrorCount: errors, - Cancelled: cancelled, - Duration: time.Since(start), - TotalRecords: total, - RecordsPerMinute: ratePerMinute(success, time.Since(start)), - } -} - -func logResult(success, errors int, startTime time.Time) { - if errors > 0 { - slog.Warn("import completed with errors", - slog.Int("success", success), - slog.Int("errors", errors)) - } - slog.Info("import completed", - slog.Int("success", success), - slog.Int("errors", errors), - slog.Duration("duration", time.Since(startTime)), - slog.String("rate", formatRate(ratePerMinute(success, time.Since(startTime))))) -} - -func backoff(attempt int) time.Duration { - if attempt <= 0 { - return BaseRetryDelay - } - - // Calculate exponential delay: BaseRetryDelay * 2^(attempt-1) - // We use uint(attempt-1) because 1<<0 is 1 (for first retry) - exp := min(attempt-1, 31) - - delay := BaseRetryDelay * time.Duration(1< MaxRetryDelay || delay <= 0 { - delay = MaxRetryDelay - } - - // Add up to 25% jitter - var jitter time.Duration - if delay > 4 { - jitter = time.Duration(rand.Int63n(int64(delay / 4))) - } - - return delay + jitter -} - -type BatchResult struct { - SuccessCount int - ErrorCount int - FailedRecords []PlayRecord - LastError error -} - -func PublishBatch(ctx context.Context, client ATProtoClient, did string, batch []PlayRecord, consumedW, consumedG int, storage cache.Storage) BatchResult { - if len(batch) == 0 { - return BatchResult{} - } - - atprotoRecords := prepareRecords(batch) - - err := client.ApplyWrites(ctx, RecordType, atprotoRecords) - if err != nil { - slog.Error("batch publish failed", logutil.Error(err)) - return BatchResult{ErrorCount: len(atprotoRecords), FailedRecords: atprotoRecords, LastError: err} - } - - logBatch(atprotoRecords, consumedW, consumedG) - - if storage != nil && did != "" { - keys := CreateRecordKeys(atprotoRecords) - cacheEntries := make(map[string][]byte) - for i, rec := range atprotoRecords { - key := keys[i] - value, _ := json.Marshal(rec) - cacheEntries[key] = value - } - if err := storage.SaveRecords(did, cacheEntries); err != nil { - slog.Debug("failed to add records to cache", logutil.Error(err)) - } - } - - return BatchResult{SuccessCount: len(atprotoRecords)} -} - -func prepareRecords(batch []PlayRecord) []PlayRecord { - atprotoRecords := make([]PlayRecord, 0, len(batch)) - for _, record := range batch { - record.Type = RecordType - record.SubmissionClientAgent = ClientAgent - atprotoRecords = append(atprotoRecords, record) - } - return atprotoRecords -} - -func logBatch(atprotoRecords []PlayRecord, consumedW, consumedG int) { - first := atprotoRecords[0] - last := atprotoRecords[len(atprotoRecords)-1] - slog.Debug("batch published", - slog.Int("records", len(atprotoRecords)), - slog.Int("writes_consumed", consumedW), - slog.Int("writes_limit", WriteLimitDay), - slog.Int("writes_remaining", WriteLimitDay-consumedW), - slog.Int("global_consumed", consumedG), - slog.Int("global_limit", GlobalLimitDay), - slog.Int("global_remaining", GlobalLimitDay-consumedG), - logutil.Track(first.TrackName, first.ArtistName(), first.PlayedTime.Time), - logutil.Track(last.TrackName, last.ArtistName(), last.PlayedTime.Time)) -} - -// removed getArtistName from here to favor the one in record.go - -func ratePerMinute(count int, duration time.Duration) float64 { - if duration == 0 { - return 0 - } - return float64(count) / duration.Minutes() -} - -type ATProtoClient interface { - ApplyWrites(ctx context.Context, collection string, records []PlayRecord) error -} - -type AuthClient interface { - GetAPIClient() *atclient.APIClient - GetDID() string -} - -type atprotoClientAdapter struct { - client *atclient.APIClient - did string -} - -func (a *atprotoClientAdapter) ApplyWrites(ctx context.Context, collection string, records []PlayRecord) error { - writes, err := prepareWrites(records, collection) - if err != nil { - return err - } - if writes == nil { - return nil - } - - return a.client.Post(ctx, syntax.NSID("com.atproto.repo.applyWrites"), map[string]any{ - "repo": a.did, - "writes": writes, - }, nil) -} diff --git a/sync/rate.go b/sync/rate.go index 105717b..fee8616 100644 --- a/sync/rate.go +++ b/sync/rate.go @@ -2,39 +2,50 @@ package sync import ( "context" + "crypto/rand" + "encoding/binary" "fmt" "log/slog" + "math" "sync" "time" + + "tangled.org/karitham.dev/lazuli/sync/logutil" ) const ( // Limits - WriteLimitDay = 9000 - GlobalLimitDay = 35000 + WriteLimitMinute = 100 + WriteLimitHour = 1000 + WriteLimitDay = 9000 + + GlobalLimitMinute = 300 + GlobalLimitHour = 3000 + GlobalLimitDay = 35000 // Costs ReadGlobalCost = 1 WriteOnlyCost = 1 WriteGlobalCost = 3 - - secondsPerDay = 86400 ) type RateLimiter interface { - // AllowBulkWrite blocks or returns error until N writes are permissible - AllowBulkWrite(ctx context.Context, n int) error - // AllowRead blocks or returns error until a read is permissible - AllowRead(ctx context.Context) error - // RefundBulkWrite restores N writes to the quota (e.g. after a failed write) - RefundBulkWrite(n int) + // AllowBulkWrite blocks or returns error until N writes are permissible. + // Returns the timestamp of when the quota was charged for bucket-accurate refunds. + AllowBulkWrite(ctx context.Context, n int) (time.Time, error) + // AllowRead blocks or returns error until a read is permissible. + AllowRead(ctx context.Context) (time.Time, error) + // RefundBulkWrite restores N writes to the quota using the original charge time. + RefundBulkWrite(ctx context.Context, n int, chargedAt time.Time) + // RefundRead restores a read to the quota using the original charge time. + RefundRead(ctx context.Context, chargedAt time.Time) // Stats returns current consumption (writes, global) - Stats() (int, int) + Stats() (int, int, error) } type KVStore interface { - Get(key string) (int, error) - Set(key string, val int) error + GetMulti(keys []string) (map[string]int, error) + IncrByMulti(counts map[string]int) error } type Clock interface { @@ -46,17 +57,19 @@ type realClock struct{} func (realClock) Now() time.Time { return time.Now().UTC() } type quotaLimiter struct { - mu sync.Mutex kv KVStore prefix string clock Clock + mu sync.Mutex } -func (l *quotaLimiter) Stats() (int, int) { - wKey, gKey := l.getKeys() - w, _ := l.kv.Get(wKey) - g, _ := l.kv.Get(gKey) - return w, g +func (l *quotaLimiter) Stats() (int, int, error) { + wd, gd, _, _, _, _ := l.getKeys(l.clock.Now()) + vals, err := l.kv.GetMulti([]string{wd, gd}) + if err != nil { + return 0, 0, fmt.Errorf("failed to get stats: %w", err) + } + return vals[wd], vals[gd], nil } func NewRateLimiter(kv KVStore) RateLimiter { @@ -67,113 +80,224 @@ func NewRateLimiter(kv KVStore) RateLimiter { } } -func (l *quotaLimiter) getKeys() (string, string) { - day := l.clock.Now().Format("2006-01-02") - return fmt.Sprintf("%s:writes:%s", l.prefix, day), fmt.Sprintf("%s:global:%s", l.prefix, day) +func (l *quotaLimiter) getKeys(t time.Time) (string, string, string, string, string, string) { + day := t.Format("2006-01-02") + hour := t.Format("2006-01-02-15") + minute := t.Format("2006-01-02-15-04") + return fmt.Sprintf("%s:writes:d:%s", l.prefix, day), fmt.Sprintf("%s:global:d:%s", l.prefix, day), + fmt.Sprintf("%s:writes:h:%s", l.prefix, hour), fmt.Sprintf("%s:global:h:%s", l.prefix, hour), + fmt.Sprintf("%s:writes:m:%s", l.prefix, minute), fmt.Sprintf("%s:global:m:%s", l.prefix, minute) } -func (l *quotaLimiter) AllowBulkWrite(ctx context.Context, n int) error { - wKey, gKey := l.getKeys() +func (l *quotaLimiter) AllowBulkWrite(ctx context.Context, n int) (time.Time, error) { wCost := n * WriteOnlyCost gCost := n * WriteGlobalCost - return l.wait(ctx, wKey, gKey, wCost, gCost, WriteLimitDay, GlobalLimitDay) -} + for { + now := l.clock.Now() + wKeys, gKeys := l.getAllKeys(now) -func (l *quotaLimiter) AllowRead(ctx context.Context) error { - _, gKey := l.getKeys() - return l.wait(ctx, "", gKey, 0, ReadGlobalCost, 0, GlobalLimitDay) -} + l.mu.Lock() + maxWait, err := l.checkQuota(now, wKeys, gKeys, wCost, gCost) + if err != nil { + l.mu.Unlock() + return now, err + } -func (l *quotaLimiter) RefundBulkWrite(n int) { - l.mu.Lock() - defer l.mu.Unlock() + if maxWait > 0 { + l.mu.Unlock() + slog.Info("Rate limit reached, sleeping until next window", slog.Duration("wait", maxWait.Round(time.Second))) + // Add a tiny bit of buffer + jitter to ensure we are definitely in the next window + wait := maxWait + 100*time.Millisecond + addJitter(100*time.Millisecond) + if err := l.sleep(ctx, wait); err != nil { + return now, err + } + continue + } - wKey, gKey := l.getKeys() - wCost := n * WriteOnlyCost - gCost := n * WriteGlobalCost + // Charge quota while holding the lock + err = l.charge(wKeys, gKeys, wCost, gCost) + l.mu.Unlock() - if currW, err := l.kv.Get(wKey); err == nil { - l.kv.Set(wKey, max(0, currW-wCost)) - } - if currG, err := l.kv.Get(gKey); err == nil { - l.kv.Set(gKey, max(0, currG-gCost)) + if err != nil { + return now, err + } + return now, nil } } -func (l *quotaLimiter) wait(ctx context.Context, wKey, gKey string, wCost, gCost, wLimit, gLimit int) error { - l.mu.Lock() - defer l.mu.Unlock() +func (l *quotaLimiter) AllowRead(ctx context.Context) (time.Time, error) { + gCost := ReadGlobalCost for { - elapsed := l.getElapsedSinceMidnight() - currW, currG := l.getCurrentConsumption(wKey, gKey) - - waitW := l.computeTargetWait(currW, wCost, wLimit, elapsed) - waitG := l.computeTargetWait(currG, gCost, gLimit, elapsed) - maxWait := max(waitG, waitW) + now := l.clock.Now() + _, gKeys := l.getAllKeys(now) - if maxWait <= 0 { - l.updateConsumption(wKey, gKey, wCost, gCost, currW, currG) - return nil + l.mu.Lock() + maxWait, err := l.checkQuota(now, nil, gKeys, 0, gCost) + if err != nil { + l.mu.Unlock() + return now, err } - if maxWait > 1*time.Minute { - slog.Info("Rate limit reached, sleeping", slog.Duration("wait", maxWait.Round(time.Second))) + if maxWait > 0 { + l.mu.Unlock() + slog.Info("Rate limit reached, sleeping until next window", slog.Duration("wait", maxWait.Round(time.Second))) + wait := maxWait + 100*time.Millisecond + addJitter(100*time.Millisecond) + if err := l.sleep(ctx, wait); err != nil { + return now, err + } + continue } + // Charge quota while holding the lock + err = l.charge(nil, gKeys, 0, gCost) l.mu.Unlock() - err := l.sleep(ctx, maxWait) - l.mu.Lock() + if err != nil { - return err + return now, err } + return now, nil } } -func (l *quotaLimiter) getElapsedSinceMidnight() float64 { - now := l.clock.Now() - midnight := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) - return now.Sub(midnight).Seconds() +func (l *quotaLimiter) getAllKeys(t time.Time) ([]string, []string) { + wd, gd, wh, gh, wm, gm := l.getKeys(t) + return []string{wm, wh, wd}, []string{gm, gh, gd} } -func (l *quotaLimiter) getCurrentConsumption(wKey, gKey string) (int, int) { - currW := 0 - if wKey != "" { - currW, _ = l.kv.Get(wKey) +func (l *quotaLimiter) RefundBulkWrite(ctx context.Context, n int, chargedAt time.Time) { + wKeys, gKeys := l.getAllKeys(chargedAt) + wCost := n * WriteOnlyCost + gCost := n * WriteGlobalCost + + l.mu.Lock() + defer l.mu.Unlock() + + updates := make(map[string]int, len(wKeys)+len(gKeys)) + for _, k := range wKeys { + updates[k] = -wCost + } + for _, k := range gKeys { + updates[k] = -gCost + } + + if err := l.kv.IncrByMulti(updates); err != nil { + slog.Error("failed to refund write quota", logutil.Error(err)) } - currG, _ := l.kv.Get(gKey) - return currW, currG } -func (l *quotaLimiter) computeTargetWait(curr, cost, limit int, elapsed float64) time.Duration { - if limit <= 0 { - return 0 +func (l *quotaLimiter) RefundRead(ctx context.Context, chargedAt time.Time) { + _, gKeys := l.getAllKeys(chargedAt) + gCost := ReadGlobalCost + + l.mu.Lock() + defer l.mu.Unlock() + + updates := make(map[string]int, len(gKeys)) + for _, k := range gKeys { + updates[k] = -gCost } - const maxCreditSeconds = 60.0 - target := float64(curr+cost) * secondsPerDay / float64(limit) - // If the target time (when this consumption is 'earned') is within the - // burst window (now + maxCreditSeconds), we don't need to wait. - if target <= elapsed+maxCreditSeconds { - return 0 + if err := l.kv.IncrByMulti(updates); err != nil { + slog.Error("failed to refund global quota (read)", logutil.Error(err)) } +} + +// checkQuota checks if the proposed cost fits within the limits. +// Returns the wait duration if over limit (0 if OK), or error. +// Must be called with lock held. +func (l *quotaLimiter) checkQuota(now time.Time, wKeys, gKeys []string, wCost, gCost int) (time.Duration, error) { + wLimits := []int{WriteLimitMinute, WriteLimitHour, WriteLimitDay} + gLimits := []int{GlobalLimitMinute, GlobalLimitHour, GlobalLimitDay} - // Otherwise, we wait until we are at the edge of the burst window. - // We cap the wait at maxCreditSeconds to avoid overly long sleeps in a single loop, - // allowing for periodic re-checks of the clock and KV store. - wait := target - (elapsed + maxCreditSeconds) - return time.Duration(min(wait, maxCreditSeconds) * float64(time.Second)) + allKeys := make([]string, 0, len(wKeys)+len(gKeys)) + allKeys = append(allKeys, wKeys...) + allKeys = append(allKeys, gKeys...) + + if len(allKeys) == 0 { + return 0, nil + } + + values, err := l.kv.GetMulti(allKeys) + if err != nil { + return 0, fmt.Errorf("failed to check quota: %w", err) + } + + maxWait := time.Duration(0) + + for i, k := range wKeys { + curr := values[k] + if curr+wCost > wLimits[i] { + wait := l.untilNextWindow(now, i) + if wait > maxWait { + maxWait = wait + } + } + } + + for i, k := range gKeys { + curr := values[k] + if curr+gCost > gLimits[i] { + wait := l.untilNextWindow(now, i) + if wait > maxWait { + maxWait = wait + } + } + } + return maxWait, nil +} + +// charge applies the cost to the keys. +// Must be called with lock held. +func (l *quotaLimiter) charge(wKeys, gKeys []string, wCost, gCost int) error { + updates := make(map[string]int, len(wKeys)+len(gKeys)) + for _, k := range wKeys { + updates[k] = wCost + } + for _, k := range gKeys { + updates[k] = gCost + } + + if len(updates) == 0 { + return nil + } + + if err := l.kv.IncrByMulti(updates); err != nil { + return fmt.Errorf("failed to charge quota: %w", err) + } + return nil } -func (l *quotaLimiter) updateConsumption(wKey, gKey string, wCost, gCost, currW, currG int) { - if wKey != "" { - l.kv.Set(wKey, currW+wCost) +func (l *quotaLimiter) untilNextWindow(now time.Time, tier int) time.Duration { + switch tier { + case 0: // Minute + return now.Truncate(time.Minute).Add(time.Minute).Sub(now) + case 1: // Hour + return now.Truncate(time.Hour).Add(time.Hour).Sub(now) + case 2: // Day + return now.Truncate(24 * time.Hour).Add(24 * time.Hour).Sub(now) + default: + // Safety fallback: if unknown tier, wait a reasonable amount (e.g. 1 minute) + // to prevent busy loops or bypassing limits. + slog.Warn("unknown rate limit tier encountered", slog.Int("tier", tier)) + return time.Minute } - l.kv.Set(gKey, currG+gCost) +} + +func addJitter(d time.Duration) time.Duration { + var b [8]byte + _, _ = rand.Read(b[:]) + n := binary.LittleEndian.Uint64(b[:]) + // Add 0-20% jitter + jitter := float64(d) * 0.2 * (float64(n) / math.MaxUint64) + return d + time.Duration(jitter) } func (l *quotaLimiter) sleep(ctx context.Context, d time.Duration) error { + if d <= 0 { + return nil + } timer := time.NewTimer(d) defer timer.Stop() select { diff --git a/sync/rate_test.go b/sync/rate_test.go index 0f1abfc..ff98821 100644 --- a/sync/rate_test.go +++ b/sync/rate_test.go @@ -2,21 +2,65 @@ package sync import ( "context" + "encoding/json" + "strings" "testing" "time" + + "github.com/bluesky-social/indigo/atproto/atclient" ) type mockKV struct { data map[string]int } +func (m *mockKV) GetMulti(keys []string) (map[string]int, error) { + out := make(map[string]int) + for _, k := range keys { + out[k] = m.data[k] + } + return out, nil +} + +func (m *mockKV) IncrByMulti(counts map[string]int) error { + for k, v := range counts { + m.data[k] += v + } + return nil +} + +// Helper for tests that inspect internal state directly func (m *mockKV) Get(key string) (int, error) { return m.data[key], nil } -func (m *mockKV) Set(key string, val int) error { - m.data[key] = val - return nil +func TestRateLimiter_Refunds(t *testing.T) { + kv := &mockKV{data: make(map[string]int)} + clock := &mockClock{now: time.Date(2026, 1, 22, 12, 0, 0, 0, time.UTC)} + limiter := "aLimiter{ + kv: kv, + prefix: "quota", + clock: clock, + } + ctx := context.Background() + + // Test Bulk Write Refund + chargedAt, _ := limiter.AllowBulkWrite(ctx, 10) + limiter.RefundBulkWrite(ctx, 10, chargedAt) + + w, g, _ := limiter.Stats() + if w != 0 || g != 0 { + t.Errorf("BulkWrite refund failed: w=%d, g=%d", w, g) + } + + // Test Read Refund + chargedAt, _ = limiter.AllowRead(ctx) + limiter.RefundRead(ctx, chargedAt) + + _, g, _ = limiter.Stats() + if g != 0 { + t.Errorf("Read refund failed: g=%d", g) + } } func TestRateLimiter_Weighting(t *testing.T) { @@ -25,21 +69,27 @@ func TestRateLimiter_Weighting(t *testing.T) { ctx := context.Background() // 1 Read = 1 Global - err := limiter.AllowRead(ctx) + _, err := limiter.AllowRead(ctx) + if err != nil { + t.Fatal(err) + } + _, g, err := limiter.Stats() if err != nil { t.Fatal(err) } - _, g := limiter.Stats() if g != 1 { t.Errorf("expected 1 global unit, got %d", g) } // 1 Write = 1 Write-Only + 3 Global - err = limiter.AllowBulkWrite(ctx, 1) + _, err = limiter.AllowBulkWrite(ctx, 1) + if err != nil { + t.Fatal(err) + } + w, g, err := limiter.Stats() if err != nil { t.Fatal(err) } - w, g := limiter.Stats() if w != 1 { t.Errorf("expected 1 write unit, got %d", w) } @@ -48,11 +98,14 @@ func TestRateLimiter_Weighting(t *testing.T) { } // Bulk Write (10 elements) = 10 Write-Only + 30 Global - err = limiter.AllowBulkWrite(ctx, 10) + _, err = limiter.AllowBulkWrite(ctx, 10) + if err != nil { + t.Fatal(err) + } + w, g, err = limiter.Stats() if err != nil { t.Fatal(err) } - w, g = limiter.Stats() if w != 11 { t.Errorf("expected 11 write units, got %d", w) } @@ -63,6 +116,7 @@ func TestRateLimiter_Weighting(t *testing.T) { func TestRateLimiter_Smoothing(t *testing.T) { kv := &mockKV{data: make(map[string]int)} + // Ensure we are at the very beginning of the minute to avoid window edge issues in test clock := &mockClock{now: time.Date(2026, 1, 22, 1, 0, 0, 0, time.UTC)} limiter := "aLimiter{ kv: kv, @@ -70,18 +124,67 @@ func TestRateLimiter_Smoothing(t *testing.T) { clock: clock, } - // 1 hour since midnight = 3600s - // allowance = (9000/86400) * 3600 = 375 - // burst allowance = (9000/86400) * (3600 + 60) = 381.25 + wd, gd, wh, gh, wm, gm := limiter.getKeys(clock.now) + kv.data[wm] = WriteLimitMinute + kv.data[wh] = 0 + kv.data[wd] = 0 + kv.data[gm] = 0 + kv.data[gh] = 0 + kv.data[gd] = 0 - _ = kv.Set("quota:writes:2026-01-22", 400) // Well over the limit + burst + // Use a context that is already cancelled to simulate what happens + // when we hit the rate limit and can't proceed. + // Ensure the store thinks we are over the limit + kv.data[wm] = WriteLimitMinute + 1 - ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) - defer cancel() + ctx, cancel := context.WithCancel(context.Background()) + cancel() - err := limiter.AllowBulkWrite(ctx, 1) - if err != context.DeadlineExceeded { - t.Errorf("expected DeadlineExceeded, got %v", err) + _, err := limiter.AllowBulkWrite(ctx, 1) + if err == nil { + t.Errorf("expected error, got nil") + } +} + +func TestRetryExhaustionMarkFailed(t *testing.T) { + ctx := context.Background() + did := "did:example:123" + storage := newMockStorage() + rec1, _ := json.Marshal(PlayRecord{TrackName: "Song 1"}) + storage.SaveRecords(did, map[string][]byte{"k1": rec1}) + + client := &mockATProtoClient{ + applyWritesFunc: func(ctx context.Context, collection string, records []PlayRecord) error { + return &atclient.APIError{StatusCode: 503} + }, + } + + oldBase := BaseRetryDelay + BaseRetryDelay = time.Nanosecond + defer func() { BaseRetryDelay = oldBase }() + + res := Publish(ctx, &mockAuthClient{did: did}, PublishOptions{ + BatchSize: 1, + ATProtoClient: client, + Storage: storage, + }) + + 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) + } + + storage.mu.Lock() + errStr, ok := storage.failed["k1"] + storage.mu.Unlock() + + if !ok { + t.Error("expected record k1 to be marked as failed") + } + if !strings.Contains(errStr, "503") { + t.Errorf("expected error string to contain 503, got %q", errStr) } } @@ -102,12 +205,15 @@ func TestRateLimiter_MidnightRollover(t *testing.T) { ctx := context.Background() // 1. Consume some quota on day 1 - err := limiter.AllowBulkWrite(ctx, 10) + _, err := limiter.AllowBulkWrite(ctx, 10) if err != nil { t.Fatal(err) } - w1, g1 := limiter.Stats() + w1, g1, err := limiter.Stats() + if err != nil { + t.Fatal(err) + } if w1 != 10 || g1 != 30 { t.Errorf("expected w=10, g=30 on day 1, got w=%d, g=%d", w1, g1) } @@ -116,24 +222,30 @@ func TestRateLimiter_MidnightRollover(t *testing.T) { clock.now = clock.now.Add(2 * time.Second) // 00:00:01 on 2026-01-23 // 3. Stats should now reflect day 2 (0) - w2, g2 := limiter.Stats() + w2, g2, err := limiter.Stats() + if err != nil { + t.Fatal(err) + } if w2 != 0 || g2 != 0 { t.Errorf("expected w=0, g=0 on day 2, got w=%d, g=%d", w2, g2) } // 4. Consumption on day 2 should not affect day 1 - err = limiter.AllowBulkWrite(ctx, 5) + _, err = limiter.AllowBulkWrite(ctx, 5) if err != nil { t.Fatal(err) } - w2, g2 = limiter.Stats() + w2, g2, err = limiter.Stats() + if err != nil { + t.Fatal(err) + } if w2 != 5 || g2 != 15 { t.Errorf("expected w=5, g=15 on day 2, got w=%d, g=%d", w2, g2) } // Verify day 1 keys are still there but not accessed by Stats() - day1WKey := "quota:writes:2026-01-22" + day1WKey := "quota:writes:d:2026-01-22" if val, _ := kv.Get(day1WKey); val != 10 { t.Errorf("expected day 1 write key to still be 10, got %d", val) } -- 2.51.2