From b0f69d3a4ffd3ab8c49bf82d20ffa1e9cfc20bed Mon Sep 17 00:00:00 2001 From: Vladimir Zagainov Date: Tue, 26 May 2026 12:15:28 +0300 Subject: [PATCH] feat: implement rate limiting and caching layer for MusicBrainz provider - Add golang.org/x/time/rate dependency for token-bucket rate limiting - Replace custom channel-based rate limiter with rate.NewLimiter(1, 1) - Add context.Context support to doGet for cancellation - Add cached_at column to external_releases via migration 005 - Implement cache hit/miss queries with TTL-based filtering - Add CacheStats type for tracking cached RGIDs - Update ExternalRelease struct with CachedAt field - Add rate limiting tests (1 req/sec enforcement, burst behavior) - Add cache tests (hit, miss, expired, mixed, empty artist) - Update migration count test for new migration --- docs/plans/2026-05-21-musicbrainz-provider.md | 14 +- go.mod | 2 + go.sum | 2 + internal/database/database.go | 17 +- internal/database/database_test.go | 6 +- internal/database/external_releases.go | 97 ++++++- internal/musicbrainz/cache.go | 53 ++++ internal/musicbrainz/cache_test.go | 255 ++++++++++++++++++ internal/musicbrainz/client.go | 91 ++----- internal/musicbrainz/client_test.go | 100 +++++-- 10 files changed, 518 insertions(+), 119 deletions(-) create mode 100644 internal/musicbrainz/cache.go create mode 100644 internal/musicbrainz/cache_test.go diff --git a/docs/plans/2026-05-21-musicbrainz-provider.md b/docs/plans/2026-05-21-musicbrainz-provider.md index d2f7d84..07944d3 100644 --- a/docs/plans/2026-05-21-musicbrainz-provider.md +++ b/docs/plans/2026-05-21-musicbrainz-provider.md @@ -56,13 +56,13 @@ Implement a MusicBrainz API provider with strict 1 request/second rate limiting, - [x] run tests - must pass before next task ### Task 2: Implement rate limiting and caching layer -- [ ] add golang.org/x/time/rate dependency to go.mod -- [ ] implement rate limiter using golang.org/x/time/rate.NewLimiter(1, 1) for 1 req/sec -- [ ] create wrapper method for rate-limited HTTP GET requests -- [ ] implement caching check: query database for existing Release Group data within TTL -- [ ] write tests for rate limiting behavior (timing tests) -- [ ] write tests for cache hit/miss logic -- [ ] run tests - must pass before next task +- [x] add golang.org/x/time/rate dependency to go.mod +- [x] implement rate limiter using golang.org/x/time/rate.NewLimiter(1, 1) for 1 req/sec +- [x] create wrapper method for rate-limited HTTP GET requests +- [x] implement caching check: query database for existing Release Group data within TTL +- [x] write tests for rate limiting behavior (timing tests) +- [x] write tests for cache hit/miss logic +- [x] run tests - must pass before next task ### Task 3: Implement MusicBrainz API endpoints and filtering - [ ] implement GetArtistReleaseGroups(artistMBID string) method diff --git a/go.mod b/go.mod index 87c8369..1847b97 100644 --- a/go.mod +++ b/go.mod @@ -7,3 +7,5 @@ require ( github.com/mattn/go-sqlite3 v1.14.22 gopkg.in/yaml.v3 v3.0.1 ) + +require golang.org/x/time v0.15.0 diff --git a/go.sum b/go.sum index 320e195..0630cee 100644 --- a/go.sum +++ b/go.sum @@ -2,6 +2,8 @@ github.com/delucks/go-subsonic v0.0.0-20240806025900-2a743ec36238 h1:uejyepOdHIS github.com/delucks/go-subsonic v0.0.0-20240806025900-2a743ec36238/go.mod h1:vnbEuj6Z20PLcHB4rrLQAOXGMjtULfMGhRVSFPcSdUo= github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU= github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= +golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= +golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= diff --git a/internal/database/database.go b/internal/database/database.go index b5260ec..bda87c8 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -115,6 +115,10 @@ func (db *DB) migrate() error { PRIMARY KEY (rgid, sent_at) );`, }, + { + name: "005_add_cached_at_to_external_releases", + sql: `ALTER TABLE external_releases ADD COLUMN cached_at DATETIME;`, + }, } for _, m := range migrations { @@ -177,12 +181,13 @@ type LocalAlbum struct { // ExternalRelease represents a row in the external_releases table. type ExternalRelease struct { - RGID string `json:"rgid"` - ArtistID string `json:"artist_id"` - Title string `json:"title"` - Type string `json:"type"` - ReleaseDate string `json:"release_date"` - IsIgnored bool `json:"is_ignored"` + RGID string `json:"rgid"` + ArtistID string `json:"artist_id"` + Title string `json:"title"` + Type string `json:"type"` + ReleaseDate string `json:"release_date"` + IsIgnored bool `json:"is_ignored"` + CachedAt time.Time `json:"cached_at"` } // NotificationSent represents a row in the notifications_sent table. diff --git a/internal/database/database_test.go b/internal/database/database_test.go index 2f3a7c3..291a979 100644 --- a/internal/database/database_test.go +++ b/internal/database/database_test.go @@ -205,8 +205,8 @@ func TestMigrationTracking(t *testing.T) { t.Fatalf("query migrations count: %v", err) } - // We have 4 recorded migrations: artist_settings, external_releases, local_albums, notifications_sent. - if count != 4 { - t.Errorf("expected 4 applied migrations, got %d", count) + // We have 5 recorded migrations: artist_settings, external_releases, local_albums, notifications_sent, cached_at column. + if count != 5 { + t.Errorf("expected 5 applied migrations, got %d", count) } } diff --git a/internal/database/external_releases.go b/internal/database/external_releases.go index 4a045ec..c0bba31 100644 --- a/internal/database/external_releases.go +++ b/internal/database/external_releases.go @@ -1,28 +1,40 @@ package database import ( + "database/sql" "fmt" + "time" + + _ "github.com/mattn/go-sqlite3" ) // GetExternalRelease retrieves an external_release row by RGID. // Returns sql.ErrNoRows if the release is not found. func GetExternalRelease(db *DB, rgid string) (*ExternalRelease, error) { var r ExternalRelease + var cachedAt sql.NullTime err := db.Conn().QueryRow( - "SELECT rgid, artist_id, title, type, release_date, is_ignored FROM external_releases WHERE rgid = ?", + "SELECT rgid, artist_id, title, type, release_date, is_ignored, cached_at FROM external_releases WHERE rgid = ?", rgid, - ).Scan(&r.RGID, &r.ArtistID, &r.Title, &r.Type, &r.ReleaseDate, &r.IsIgnored) + ).Scan(&r.RGID, &r.ArtistID, &r.Title, &r.Type, &r.ReleaseDate, &r.IsIgnored, &cachedAt) if err != nil { return nil, err } + if cachedAt.Valid { + r.CachedAt = cachedAt.Time + } return &r, nil } // SaveExternalRelease inserts or replaces an external_release row. func SaveExternalRelease(db *DB, release *ExternalRelease) error { + cachedAtStr := "" + if !release.CachedAt.IsZero() { + cachedAtStr = release.CachedAt.Format("2006-01-02 15:04:05") + } _, err := db.Conn().Exec( - "INSERT OR REPLACE INTO external_releases (rgid, artist_id, title, type, release_date, is_ignored) VALUES (?, ?, ?, ?, ?, ?)", - release.RGID, release.ArtistID, release.Title, release.Type, release.ReleaseDate, release.IsIgnored, + "INSERT OR REPLACE INTO external_releases (rgid, artist_id, title, type, release_date, is_ignored, cached_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + release.RGID, release.ArtistID, release.Title, release.Type, release.ReleaseDate, release.IsIgnored, cachedAtStr, ) if err != nil { return fmt.Errorf("save external release: %w", err) @@ -33,7 +45,7 @@ func SaveExternalRelease(db *DB, release *ExternalRelease) error { // GetExternalReleasesByArtist returns all external_release rows for a given artist_id. func GetExternalReleasesByArtist(db *DB, artistID string) ([]ExternalRelease, error) { rows, err := db.Conn().Query( - "SELECT rgid, artist_id, title, type, release_date, is_ignored FROM external_releases WHERE artist_id = ?", + "SELECT rgid, artist_id, title, type, release_date, is_ignored, cached_at FROM external_releases WHERE artist_id = ?", artistID, ) if err != nil { @@ -44,9 +56,13 @@ func GetExternalReleasesByArtist(db *DB, artistID string) ([]ExternalRelease, er var results []ExternalRelease for rows.Next() { var r ExternalRelease - if err := rows.Scan(&r.RGID, &r.ArtistID, &r.Title, &r.Type, &r.ReleaseDate, &r.IsIgnored); err != nil { + var cachedAt sql.NullTime + if err := rows.Scan(&r.RGID, &r.ArtistID, &r.Title, &r.Type, &r.ReleaseDate, &r.IsIgnored, &cachedAt); err != nil { return nil, fmt.Errorf("scan external release: %w", err) } + if cachedAt.Valid { + r.CachedAt = cachedAt.Time + } results = append(results, r) } if err := rows.Err(); err != nil { @@ -58,7 +74,7 @@ func GetExternalReleasesByArtist(db *DB, artistID string) ([]ExternalRelease, er // GetIgnoredReleases returns all external_release rows where is_ignored = 1. func GetIgnoredReleases(db *DB) ([]ExternalRelease, error) { rows, err := db.Conn().Query( - "SELECT rgid, artist_id, title, type, release_date, is_ignored FROM external_releases WHERE is_ignored = 1", + "SELECT rgid, artist_id, title, type, release_date, is_ignored, cached_at FROM external_releases WHERE is_ignored = 1", ) if err != nil { return nil, fmt.Errorf("query ignored releases: %w", err) @@ -68,9 +84,13 @@ func GetIgnoredReleases(db *DB) ([]ExternalRelease, error) { var results []ExternalRelease for rows.Next() { var r ExternalRelease - if err := rows.Scan(&r.RGID, &r.ArtistID, &r.Title, &r.Type, &r.ReleaseDate, &r.IsIgnored); err != nil { + var cachedAt sql.NullTime + if err := rows.Scan(&r.RGID, &r.ArtistID, &r.Title, &r.Type, &r.ReleaseDate, &r.IsIgnored, &cachedAt); err != nil { return nil, fmt.Errorf("scan ignored release: %w", err) } + if cachedAt.Valid { + r.CachedAt = cachedAt.Time + } results = append(results, r) } if err := rows.Err(); err != nil { @@ -99,3 +119,64 @@ func SetReleaseIgnored(db *DB, rgid string, ignored bool) error { return nil } + +// DeleteExternalReleasesByArtist removes all external_release rows for a given artist_id. +func DeleteExternalReleasesByArtist(db *DB, artistID string) error { + _, err := db.Conn().Exec( + "DELETE FROM external_releases WHERE artist_id = ?", + artistID, + ) + if err != nil { + return fmt.Errorf("delete external releases by artist: %w", err) + } + return nil +} + +// CountExternalReleases returns the total number of external_release rows. +func CountExternalReleases(db *DB) (int, error) { + var count int + err := db.Conn().QueryRow("SELECT COUNT(*) FROM external_releases").Scan(&count) + if err != nil { + return 0, fmt.Errorf("count external releases: %w", err) + } + return count, nil +} + +// GetExternalReleasesByArtistWithCache returns cached external_release rows for a given artist_id +// that are within the specified TTL. +func GetExternalReleasesByArtistWithCache(db *DB, artistID string, ttl time.Duration) ([]ExternalRelease, error) { + cutoff := time.Now().Add(-ttl) + rows, err := db.Conn().Query( + "SELECT rgid, artist_id, title, type, release_date, is_ignored, cached_at FROM external_releases WHERE artist_id = ? AND cached_at >= ?", + artistID, cutoff.Format("2006-01-02 15:04:05"), + ) + if err != nil { + return nil, fmt.Errorf("query cached external releases: %w", err) + } + defer rows.Close() + + var results []ExternalRelease + for rows.Next() { + var r ExternalRelease + var releaseDate sql.NullString + var releaseType sql.NullString + var cachedAt sql.NullTime + if err := rows.Scan(&r.RGID, &r.ArtistID, &r.Title, &releaseType, &releaseDate, &r.IsIgnored, &cachedAt); err != nil { + return nil, fmt.Errorf("scan cached external release: %w", err) + } + if releaseType.Valid { + r.Type = releaseType.String + } + if releaseDate.Valid { + r.ReleaseDate = releaseDate.String + } + if cachedAt.Valid { + r.CachedAt = cachedAt.Time + } + results = append(results, r) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate cached external releases: %w", err) + } + return results, nil +} diff --git a/internal/musicbrainz/cache.go b/internal/musicbrainz/cache.go new file mode 100644 index 0000000..a88b0a4 --- /dev/null +++ b/internal/musicbrainz/cache.go @@ -0,0 +1,53 @@ +package musicbrainz + +import ( + "fmt" + "time" + + "naviwatcher/internal/database" +) + +// CacheStats holds the result of a cache lookup for a given artist. +type CacheStats struct { + // CachedRGIDs is the list of RGIDs that are currently cached (within TTL). + CachedRGIDs []string + // CacheHitCount is the number of entries found in cache. + CacheHitCount int +} + +// IsCached returns true if the given RGID is in the cached set. +func (cs *CacheStats) IsCached(rgid string) bool { + for _, id := range cs.CachedRGIDs { + if id == rgid { + return true + } + } + return false +} + +// GetCachedReleases queries the external_releases table for entries +// belonging to the given artist that were cached within the specified TTL. +// It returns a CacheStats with the list of valid RGIDs already in cache. +func GetCachedReleases(db *database.DB, artistID string, ttl time.Duration) (*CacheStats, error) { + releases, err := database.GetExternalReleasesByArtistWithCache(db, artistID, ttl) + if err != nil { + return nil, fmt.Errorf("get cached releases: %w", err) + } + + stats := &CacheStats{} + for _, r := range releases { + stats.CachedRGIDs = append(stats.CachedRGIDs, r.RGID) + stats.CacheHitCount++ + } + return stats, nil +} + +// IsArtistCacheValid checks whether the cache for an artist is still valid. +// Returns true if any entries exist within the TTL for this artist. +func IsArtistCacheValid(db *database.DB, artistID string, ttl time.Duration) (bool, error) { + stats, err := GetCachedReleases(db, artistID, ttl) + if err != nil { + return false, err + } + return stats.CacheHitCount > 0, nil +} diff --git a/internal/musicbrainz/cache_test.go b/internal/musicbrainz/cache_test.go new file mode 100644 index 0000000..ea51383 --- /dev/null +++ b/internal/musicbrainz/cache_test.go @@ -0,0 +1,255 @@ +package musicbrainz + +import ( + "testing" + "time" + + "naviwatcher/internal/database" +) + +func insertTestArtistForCache(db *database.DB, id string) error { + _, err := db.Conn().Exec( + "INSERT OR IGNORE INTO artist_settings (id, name) VALUES (?, ?)", + id, "Test Artist "+id, + ) + return err +} + +func TestGetCachedReleases_CacheHit(t *testing.T) { + db, err := database.New(":memory:") + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + artistID := "artist-cache-hit" + if err := insertTestArtistForCache(db, artistID); err != nil { + t.Fatalf("insertTestArtist: %v", err) + } + + // Insert releases with recent cached_at timestamps + now := time.Now().Format("2006-01-02 15:04:05") + _, err = db.Conn().Exec( + "INSERT INTO external_releases (rgid, artist_id, title, type, cached_at) VALUES (?, ?, ?, ?, ?)", + "rg-hit-1", artistID, "Cached Album 1", "album", now, + ) + if err != nil { + t.Fatalf("insert rg-hit-1: %v", err) + } + _, err = db.Conn().Exec( + "INSERT INTO external_releases (rgid, artist_id, title, type, cached_at) VALUES (?, ?, ?, ?, ?)", + "rg-hit-2", artistID, "Cached Album 2", "single", now, + ) + if err != nil { + t.Fatalf("insert rg-hit-2: %v", err) + } + + ttl := 24 * time.Hour + stats, err := GetCachedReleases(db, artistID, ttl) + if err != nil { + t.Fatalf("GetCachedReleases() error: %v", err) + } + + if stats.CacheHitCount != 2 { + t.Errorf("CacheHitCount = %d, want 2", stats.CacheHitCount) + } + + if !stats.IsCached("rg-hit-1") { + t.Error("expected rg-hit-1 to be cached") + } + if !stats.IsCached("rg-hit-2") { + t.Error("expected rg-hit-2 to be cached") + } +} + +func TestGetCachedReleases_CacheMiss_Expired(t *testing.T) { + db, err := database.New(":memory:") + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + artistID := "artist-cache-miss" + if err := insertTestArtistForCache(db, artistID); err != nil { + t.Fatalf("insertTestArtist: %v", err) + } + + // Insert a release with an expired cached_at (48 hours ago) + expired := time.Now().Add(-48 * time.Hour).Format("2006-01-02 15:04:05") + _, err = db.Conn().Exec( + "INSERT INTO external_releases (rgid, artist_id, title, type, cached_at) VALUES (?, ?, ?, ?, ?)", + "rg-expired", artistID, "Expired Album", "album", expired, + ) + if err != nil { + t.Fatalf("insert expired release: %v", err) + } + + ttl := 24 * time.Hour + stats, err := GetCachedReleases(db, artistID, ttl) + if err != nil { + t.Fatalf("GetCachedReleases() error: %v", err) + } + + if stats.CacheHitCount != 0 { + t.Errorf("CacheHitCount = %d, want 0 (expired entry should not be cached)", stats.CacheHitCount) + } + + if stats.IsCached("rg-expired") { + t.Error("expected rg-expired to NOT be cached") + } +} + +func TestGetCachedReleases_CacheMiss_NoCachedAt(t *testing.T) { + db, err := database.New(":memory:") + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + artistID := "artist-no-cached" + if err := insertTestArtistForCache(db, artistID); err != nil { + t.Fatalf("insertTestArtist: %v", err) + } + + // Insert a release WITHOUT cached_at (NULL) + _, err = db.Conn().Exec( + "INSERT INTO external_releases (rgid, artist_id, title, type) VALUES (?, ?, ?, ?)", + "rg-nocached", artistID, "Uncached Album", "album", + ) + if err != nil { + t.Fatalf("insert uncached release: %v", err) + } + + ttl := 24 * time.Hour + stats, err := GetCachedReleases(db, artistID, ttl) + if err != nil { + t.Fatalf("GetCachedReleases() error: %v", err) + } + + if stats.CacheHitCount != 0 { + t.Errorf("CacheHitCount = %d, want 0 (NULL cached_at should not be cached)", stats.CacheHitCount) + } +} + +func TestGetCachedReleases_EmptyArtist(t *testing.T) { + db, err := database.New(":memory:") + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + ttl := 24 * time.Hour + stats, err := GetCachedReleases(db, "nonexistent-artist", ttl) + if err != nil { + t.Fatalf("GetCachedReleases() error: %v", err) + } + + if stats.CacheHitCount != 0 { + t.Errorf("CacheHitCount = %d, want 0 for nonexistent artist", stats.CacheHitCount) + } +} + +func TestIsArtistCacheValid_Valid(t *testing.T) { + db, err := database.New(":memory:") + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + artistID := "artist-valid-cache" + if err := insertTestArtistForCache(db, artistID); err != nil { + t.Fatalf("insertTestArtist: %v", err) + } + + now := time.Now().Format("2006-01-02 15:04:05") + _, err = db.Conn().Exec( + "INSERT INTO external_releases (rgid, artist_id, title, cached_at) VALUES (?, ?, ?, ?)", + "rg-valid", artistID, "Valid Album", now, + ) + if err != nil { + t.Fatalf("insert: %v", err) + } + + ttl := 24 * time.Hour + valid, err := IsArtistCacheValid(db, artistID, ttl) + if err != nil { + t.Fatalf("IsArtistCacheValid() error: %v", err) + } + if !valid { + t.Error("expected cache to be valid") + } +} + +func TestIsArtistCacheValid_Invalid(t *testing.T) { + db, err := database.New(":memory:") + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + ttl := 24 * time.Hour + + // No releases at all + valid, err := IsArtistCacheValid(db, "no-releases", ttl) + if err != nil { + t.Fatalf("IsArtistCacheValid() error: %v", err) + } + if valid { + t.Error("expected cache to be invalid for artist with no releases") + } +} + +func TestCacheStats_IsCached_Empty(t *testing.T) { + stats := &CacheStats{} + if stats.IsCached("anything") { + t.Error("expected IsCached to return false for empty stats") + } +} + +func TestGetCachedReleases_MixedExpiry(t *testing.T) { + db, err := database.New(":memory:") + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + artistID := "artist-mixed" + if err := insertTestArtistForCache(db, artistID); err != nil { + t.Fatalf("insertTestArtist: %v", err) + } + + now := time.Now().Format("2006-01-02 15:04:05") + expired := time.Now().Add(-48 * time.Hour).Format("2006-01-02 15:04:05") + + // Mix of fresh and expired + _, err = db.Conn().Exec( + "INSERT INTO external_releases (rgid, artist_id, title, cached_at) VALUES (?, ?, ?, ?)", + "rg-fresh", artistID, "Fresh Album", now, + ) + if err != nil { + t.Fatalf("insert fresh: %v", err) + } + _, err = db.Conn().Exec( + "INSERT INTO external_releases (rgid, artist_id, title, cached_at) VALUES (?, ?, ?, ?)", + "rg-old", artistID, "Old Album", expired, + ) + if err != nil { + t.Fatalf("insert old: %v", err) + } + + ttl := 24 * time.Hour + stats, err := GetCachedReleases(db, artistID, ttl) + if err != nil { + t.Fatalf("GetCachedReleases() error: %v", err) + } + + if stats.CacheHitCount != 1 { + t.Errorf("CacheHitCount = %d, want 1 (only fresh entry)", stats.CacheHitCount) + } + if !stats.IsCached("rg-fresh") { + t.Error("expected rg-fresh to be cached") + } + if stats.IsCached("rg-old") { + t.Error("expected rg-old to NOT be cached (expired)") + } +} diff --git a/internal/musicbrainz/client.go b/internal/musicbrainz/client.go index 1446c3f..36b60df 100644 --- a/internal/musicbrainz/client.go +++ b/internal/musicbrainz/client.go @@ -1,12 +1,14 @@ package musicbrainz import ( + "context" "encoding/xml" "fmt" "io" "net/http" "time" + "golang.org/x/time/rate" "naviwatcher/internal/config" ) @@ -16,77 +18,26 @@ type MusicBrainzClient struct { httpClient *http.Client userAgent string baseURL string - rateLimiter *rateLimiter -} - -// rateLimiter wraps a token-bucket rate limiter for API calls. -type rateLimiter struct { - // tokens is a channel-based semaphore for rate limiting. - // It is filled at a fixed interval by a background goroutine. - tokens chan struct{} - done chan struct{} -} - -// newRateLimiter creates a rate limiter that allows maxCalls per second. -// It immediately fills the bucket and starts a refill goroutine. -func newRateLimiter(callsPerSecond int) *rateLimiter { - rl := &rateLimiter{ - tokens: make(chan struct{}, callsPerSecond), - done: make(chan struct{}), - } - // Fill the bucket initially - for i := 0; i < callsPerSecond; i++ { - rl.tokens <- struct{}{} - } - // Refill at the specified interval - interval := time.Second / time.Duration(callsPerSecond) - go func() { - ticker := time.NewTicker(interval) - defer ticker.Stop() - for { - select { - case <-ticker.C: - select { - case rl.tokens <- struct{}{}: - default: - // bucket full, skip - } - case <-rl.done: - return - } - } - }() - return rl -} - -// wait blocks until a token is available or the rate limiter is stopped. -func (rl *rateLimiter) wait() { - <-rl.tokens -} - -// stop terminates the refill goroutine. -func (rl *rateLimiter) stop() { - close(rl.done) + rateLimiter *rate.Limiter } // NewClient creates a new MusicBrainzClient from the given configuration. // It initializes the HTTP client with a 30-second timeout and sets up // a rate limiter for 1 request per second as required by MusicBrainz policy. func NewClient(cfg config.MusicBrainzConfig) *MusicBrainzClient { - rl := newRateLimiter(1) // 1 request per second return &MusicBrainzClient{ httpClient: &http.Client{ Timeout: 30 * time.Second, }, userAgent: cfg.UserAgent, baseURL: "https://musicbrainz.org/ws/2", - rateLimiter: rl, + rateLimiter: rate.NewLimiter(rate.Limit(1), 1), } } // NewClientWithLimiter creates a MusicBrainzClient with a custom rate limiter. // This is primarily used for testing to inject a mock rate limiter. -func NewClientWithLimiter(cfg config.MusicBrainzConfig, rl *rateLimiter) *MusicBrainzClient { +func NewClientWithLimiter(cfg config.MusicBrainzConfig, rl *rate.Limiter) *MusicBrainzClient { return &MusicBrainzClient{ httpClient: &http.Client{ Timeout: 30 * time.Second, @@ -97,17 +48,19 @@ func NewClientWithLimiter(cfg config.MusicBrainzConfig, rl *rateLimiter) *MusicB } } -// Close cleans up the rate limiter goroutine. -func (c *MusicBrainzClient) Close() { - c.rateLimiter.stop() -} +// Close is a no-op for the x/time/rate-based client (the limiter does not +// spawn goroutines), but retained for API compatibility. +func (c *MusicBrainzClient) Close() {} // doGet performs a rate-limited HTTP GET request to the MusicBrainz API. -// It sets the proper User-Agent header and returns the response body. -func (c *MusicBrainzClient) doGet(path string) ([]byte, error) { - c.rateLimiter.wait() +// It blocks until the rate limiter allows the request, then sets the proper +// User-Agent header and returns the response body. +func (c *MusicBrainzClient) doGet(ctx context.Context, path string) ([]byte, error) { + if err := c.rateLimiter.Wait(ctx); err != nil { + return nil, fmt.Errorf("rate limiter wait: %w", err) + } - req, err := http.NewRequest(http.MethodGet, c.baseURL+path, nil) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+path, nil) if err != nil { return nil, fmt.Errorf("create request: %w", err) } @@ -152,12 +105,12 @@ type mbArtistCredit struct { // mbReleaseGroup represents the XML structure of a single release-group // in the MusicBrainz release-group list response. type mbReleaseGroup struct { - ID string `xml:"id,attr"` - Title string `xml:"title"` - Type string `xml:"type,attr"` - Status string `xml:"status,attr"` - ArtistCredit mbArtistCredit `xml:"artist-credit"` - ReleaseDate string `xml:"first-release-date"` + ID string `xml:"id,attr"` + Title string `xml:"title"` + Type string `xml:"type,attr"` + Status string `xml:"status,attr"` + ArtistCredit mbArtistCredit `xml:"artist-credit"` + ReleaseDate string `xml:"first-release-date"` } // mbReleaseGroupListXML wraps the release-group-list element to properly @@ -169,7 +122,7 @@ type mbReleaseGroupListXML struct { // mbReleaseGroupList represents the XML structure of a release-group list response. type mbReleaseGroupList struct { - XMLName xml.Name `xml:"metadata"` + XMLName xml.Name `xml:"metadata"` ReleaseGroupList mbReleaseGroupListXML `xml:"release-group-list"` } diff --git a/internal/musicbrainz/client_test.go b/internal/musicbrainz/client_test.go index 6c1e8f3..885a6d2 100644 --- a/internal/musicbrainz/client_test.go +++ b/internal/musicbrainz/client_test.go @@ -1,18 +1,16 @@ package musicbrainz import ( + "context" "net/http" "net/http/httptest" "testing" "time" + "golang.org/x/time/rate" "naviwatcher/internal/config" ) -func newUnbufferedRateLimiter() *rateLimiter { - return newRateLimiter(1000) // high rate to avoid blocking in tests -} - func TestNewClient_ValidConfig(t *testing.T) { cfg := config.MusicBrainzConfig{ UserAgent: "naviwatcher/0.1.0 (test@example.com)", @@ -48,7 +46,7 @@ func TestNewClientWithLimiter(t *testing.T) { UserAgent: "naviwatcher/0.1.0 (test@example.com)", } - rl := newRateLimiter(1) + rl := rate.NewLimiter(rate.Limit(1), 1) client := NewClientWithLimiter(cfg, rl) defer client.Close() @@ -78,7 +76,7 @@ func TestDoGet_Success(t *testing.T) { UserAgent: "naviwatcher/0.1.0 (test@example.com)", } - rl := newUnbufferedRateLimiter() + rl := rate.NewLimiter(rate.Limit(1000), 1000) // high rate to avoid blocking in tests client := &MusicBrainzClient{ httpClient: server.Client(), userAgent: cfg.UserAgent, @@ -87,7 +85,7 @@ func TestDoGet_Success(t *testing.T) { } defer client.Close() - body, err := client.doGet("/test") + body, err := client.doGet(context.Background(), "/test") if err != nil { t.Fatalf("doGet() error = %v", err) } @@ -108,7 +106,7 @@ func TestDoGet_Non200Status(t *testing.T) { UserAgent: "naviwatcher/0.1.0 (test@example.com)", } - rl := newUnbufferedRateLimiter() + rl := rate.NewLimiter(rate.Limit(1000), 1000) client := &MusicBrainzClient{ httpClient: server.Client(), userAgent: cfg.UserAgent, @@ -117,7 +115,7 @@ func TestDoGet_Non200Status(t *testing.T) { } defer client.Close() - _, err := client.doGet("/test") + _, err := client.doGet(context.Background(), "/test") if err == nil { t.Fatal("doGet() expected error for non-200 status, got nil") } @@ -131,7 +129,7 @@ func TestDoGet_ServerUnreachable(t *testing.T) { UserAgent: "naviwatcher/0.1.0 (test@example.com)", } - rl := newUnbufferedRateLimiter() + rl := rate.NewLimiter(rate.Limit(1000), 1000) client := &MusicBrainzClient{ httpClient: server.Client(), userAgent: cfg.UserAgent, @@ -140,28 +138,78 @@ func TestDoGet_ServerUnreachable(t *testing.T) { } defer client.Close() - _, err := client.doGet("/test") + _, err := client.doGet(context.Background(), "/test") if err == nil { t.Fatal("doGet() expected error for unreachable server, got nil") } } -func TestRateLimiter_BasicBehavior(t *testing.T) { - // Test that rate limiter can produce tokens - rl := newRateLimiter(1) - defer rl.stop() +func TestDoGet_ContextCancellation(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Write([]byte(`ok`)) + })) + defer server.Close() - // Should be able to get a token immediately (bucket was pre-filled) - done := make(chan struct{}) - go func() { - rl.wait() - close(done) - }() + cfg := config.MusicBrainzConfig{ + UserAgent: "naviwatcher/0.1.0 (test@example.com)", + } - select { - case <-done: - // success - case <-time.After(2 * time.Second): - t.Fatal("rateLimiter.wait() blocked on pre-filled bucket") + // Use a rate limiter with 0 burst to force blocking on Wait + rl := rate.NewLimiter(rate.Limit(0), 0) + client := &MusicBrainzClient{ + httpClient: server.Client(), + userAgent: cfg.UserAgent, + baseURL: server.URL, + rateLimiter: rl, + } + defer client.Close() + + ctx, cancel := context.WithCancel(context.Background()) + cancel() // cancel immediately + + _, err := client.doGet(ctx, "/test") + if err == nil { + t.Fatal("doGet() expected error for cancelled context, got nil") + } +} + +func TestRateLimiter_OnePerSecond(t *testing.T) { + // Verify that the rate limiter enforces approximately 1 request per second + rl := rate.NewLimiter(rate.Limit(1), 1) + + // First request should be immediate (burst of 1) + start := time.Now() + if err := rl.Wait(context.Background()); err != nil { + t.Fatalf("first Wait() error: %v", err) + } + elapsed := time.Since(start) + if elapsed > 100*time.Millisecond { + t.Errorf("first Wait() took %v, expected near-instant", elapsed) + } + + // Second request should block for approximately 1 second + start = time.Now() + if err := rl.Wait(context.Background()); err != nil { + t.Fatalf("second Wait() error: %v", err) + } + elapsed = time.Since(start) + if elapsed < 800*time.Millisecond { + t.Errorf("second Wait() took %v, expected at least ~1s", elapsed) + } + if elapsed > 2*time.Second { + t.Errorf("second Wait() took %v, expected less than 2s", elapsed) + } +} + +func TestRateLimiter_BurstBehavior(t *testing.T) { + // With burst=1, the first request should be immediate + rl := rate.NewLimiter(rate.Limit(1), 1) + + start := time.Now() + rl.Wait(context.Background()) + elapsed := time.Since(start) + + if elapsed > 50*time.Millisecond { + t.Errorf("burst Wait() took %v, expected near-instant", elapsed) } }