diff --git a/internal/database/database.go b/internal/database/database.go index d7a7be4..9154cfe 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -980,23 +980,20 @@ func (d *Database) GetStatsContext(ctx context.Context) (Stats, error) { if err != nil { return stats, fmt.Errorf("failed to count IPv6 routes: %w", err) } + stats.IPv4Routes = v4Count + stats.IPv6Routes = v6Count stats.LiveRoutes = v4Count + v6Count - // Get oldest and newest route timestamps - routeTimestampQuery := ` - SELECT MIN(last_updated), MAX(last_updated) FROM ( - SELECT last_updated FROM live_routes_v4 - UNION ALL - SELECT last_updated FROM live_routes_v6 - ) - ` - var oldestRoute, newestRoute *time.Time - err = d.db.QueryRowContext(ctx, routeTimestampQuery).Scan(&oldestRoute, &newestRoute) + // Get oldest and newest route timestamps. Each query reads a single row from + // one end of the last_updated index, so the cost is a log-time index lookup + // rather than a full scan of both route tables. Selecting the last_updated + // column directly (rather than MIN/MAX, whose result has no column type) lets + // the driver parse the DATETIME value into time.Time; the union scan aggregate + // used before returned an untyped string and logged a warning on every call. + stats.OldestRoute, stats.NewestRoute, err = d.routeTimestampRange(ctx) if err != nil { + // Display-only fields; log but keep the rest of the stats. d.logger.Warn("Failed to get route timestamps", "error", err) - } else { - stats.OldestRoute = oldestRoute - stats.NewestRoute = newestRoute } // Get prefix distribution @@ -1009,6 +1006,64 @@ func (d *Database) GetStatsContext(ctx context.Context) (Stats, error) { return stats, nil } +// routeTimestampRange returns the earliest and latest last_updated across both +// live route tables, or nil values when both tables are empty. Each query reads +// one row from an end of the last_updated index rather than scanning the tables. +func (d *Database) routeTimestampRange(ctx context.Context) (oldest, newest *time.Time, err error) { + oldestV4, ok, err := d.scanRouteTimestamp(ctx, + "SELECT last_updated FROM live_routes_v4 ORDER BY last_updated ASC LIMIT 1") + if err != nil { + return nil, nil, err + } + if ok { + oldest = &oldestV4 + } + + oldestV6, ok, err := d.scanRouteTimestamp(ctx, + "SELECT last_updated FROM live_routes_v6 ORDER BY last_updated ASC LIMIT 1") + if err != nil { + return nil, nil, err + } + if ok && (oldest == nil || oldestV6.Before(*oldest)) { + oldest = &oldestV6 + } + + newestV4, ok, err := d.scanRouteTimestamp(ctx, + "SELECT last_updated FROM live_routes_v4 ORDER BY last_updated DESC LIMIT 1") + if err != nil { + return nil, nil, err + } + if ok { + newest = &newestV4 + } + + newestV6, ok, err := d.scanRouteTimestamp(ctx, + "SELECT last_updated FROM live_routes_v6 ORDER BY last_updated DESC LIMIT 1") + if err != nil { + return nil, nil, err + } + if ok && (newest == nil || newestV6.After(*newest)) { + newest = &newestV6 + } + + return oldest, newest, nil +} + +// scanRouteTimestamp runs a single-row timestamp query. ok is false when the +// table is empty. The query selects the last_updated column directly so the +// driver parses the DATETIME value into a time.Time. +func (d *Database) scanRouteTimestamp(ctx context.Context, query string) (ts time.Time, ok bool, err error) { + err = d.db.QueryRowContext(ctx, query).Scan(&ts) + switch { + case errors.Is(err, sql.ErrNoRows): + return time.Time{}, false, nil + case err != nil: + return time.Time{}, false, err + default: + return ts, true, nil + } +} + // UpsertLiveRoute inserts or updates a live route func (d *Database) UpsertLiveRoute(route *LiveRoute) error { d.lock("UpsertLiveRoute") diff --git a/internal/database/database_test.go b/internal/database/database_test.go index fa4bb48..e9d3c0b 100644 --- a/internal/database/database_test.go +++ b/internal/database/database_test.go @@ -10,6 +10,7 @@ import ( "git.eeqj.de/sneak/routewatch/internal/config" "git.eeqj.de/sneak/routewatch/internal/logger" + "github.com/google/uuid" ) // tempStoreMemory is the PRAGMA temp_store value meaning "hold temp B-trees in @@ -425,6 +426,105 @@ func TestBatchWriteDuringCheckpoint(t *testing.T) { wg.Wait() } +// TestStatsRouteTimestampsAndCounts checks GetStatsContext reports the correct +// route counts and the oldest/newest last_updated across both route tables. The +// old union-scan query read the aggregate result into *time.Time, which the +// driver could not parse, so it logged a warning every call and left both +// timestamps nil; this asserts they are populated from the right rows. +func TestStatsRouteTimestampsAndCounts(t *testing.T) { + cfg := &config.Config{StateDir: t.TempDir()} + + db, err := New(cfg, logger.New()) + if err != nil { + t.Fatalf("failed to create database: %v", err) + } + defer func() { _ = db.Close() }() + + ctx := context.Background() + + // Empty database: no routes, so both timestamps are nil and no error. + empty, err := db.GetStatsContext(ctx) + if err != nil { + t.Fatalf("GetStatsContext on empty database: %v", err) + } + if empty.OldestRoute != nil || empty.NewestRoute != nil { + t.Fatalf("empty database timestamps = (%v, %v), want (nil, nil)", + empty.OldestRoute, empty.NewestRoute) + } + if empty.LiveRoutes != 0 { + t.Fatalf("empty database LiveRoutes = %d, want 0", empty.LiveRoutes) + } + + base := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) + oldest := base + middle := base.Add(time.Minute) + newest := base.Add(2 * time.Minute) + + mkV4 := func(prefix string, asn int, ts time.Time) *LiveRoute { + start, end, rerr := CalculateIPv4Range(prefix) + if rerr != nil { + t.Fatalf("CalculateIPv4Range(%s): %v", prefix, rerr) + } + + return &LiveRoute{ + ID: uuid.New(), + Prefix: prefix, + MaskLength: 24, + IPVersion: ipVersionV4, + OriginASN: asn, + PeerIP: "192.0.2.1", + ASPath: []int{asn}, + NextHop: "192.0.2.254", + LastUpdated: ts, + V4IPStart: &start, + V4IPEnd: &end, + } + } + + // Two IPv4 routes (one oldest, one middle) and one IPv6 route (newest). + routes := []*LiveRoute{ + mkV4("198.51.100.0/24", 64500, middle), + mkV4("203.0.113.0/24", 64501, oldest), + { + ID: uuid.New(), + Prefix: "2001:db8::/32", + MaskLength: 32, + IPVersion: ipVersionV6, + OriginASN: 64502, + PeerIP: "2001:db8::1", + ASPath: []int{64502}, + NextHop: "2001:db8::ffff", + LastUpdated: newest, + }, + } + for _, route := range routes { + if err := db.UpsertLiveRoute(route); err != nil { + t.Fatalf("UpsertLiveRoute(%s): %v", route.Prefix, err) + } + } + + stats, err := db.GetStatsContext(ctx) + if err != nil { + t.Fatalf("GetStatsContext: %v", err) + } + + if stats.IPv4Routes != 2 { + t.Errorf("IPv4Routes = %d, want 2", stats.IPv4Routes) + } + if stats.IPv6Routes != 1 { + t.Errorf("IPv6Routes = %d, want 1", stats.IPv6Routes) + } + if stats.LiveRoutes != 3 { + t.Errorf("LiveRoutes = %d, want 3", stats.LiveRoutes) + } + if stats.OldestRoute == nil || !stats.OldestRoute.Equal(oldest) { + t.Errorf("OldestRoute = %v, want %v", stats.OldestRoute, oldest) + } + if stats.NewestRoute == nil || !stats.NewestRoute.Equal(newest) { + t.Errorf("NewestRoute = %v, want %v", stats.NewestRoute, newest) + } +} + func BenchmarkIPToUint32(b *testing.B) { ip := net.ParseIP("192.168.1.1") b.ResetTimer() diff --git a/internal/database/interface.go b/internal/database/interface.go index 7f3f8e9..038dfd5 100644 --- a/internal/database/interface.go +++ b/internal/database/interface.go @@ -18,6 +18,8 @@ type Stats struct { Peers int FileSizeBytes int64 LiveRoutes int + IPv4Routes int + IPv6Routes int OldestRoute *time.Time NewestRoute *time.Time IPv4PrefixDistribution []PrefixDistribution diff --git a/internal/server/handlers.go b/internal/server/handlers.go index d04d0fd..1d0d577 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -179,37 +179,14 @@ func (s *Server) handleStatusJSON() http.HandlerFunc { metrics := s.streamer.GetMetrics() - // Get database stats with timeout. The channels are buffered so the - // goroutine's send never blocks if the timeout wins and nothing here - // receives; otherwise it would block forever and leak. - statsChan := make(chan database.Stats, 1) - errChan := make(chan error, 1) - - go func() { - dbStats, err := s.db.GetStatsContext(ctx) - if err != nil { - s.logger.Debug("Database stats query failed", "error", err) - errChan <- err - - return - } - statsChan <- dbStats - }() - - var dbStats database.Stats - select { - case <-ctx.Done(): - s.logger.Error("Database stats timeout in status.json") - writeJSONError(w, http.StatusRequestTimeout, "Database timeout") - - return - case err := <-errChan: + // Serve database statistics from the cache, which runs the table scans at + // most once per interval so this request does not. + dbStats, err := s.stats.get() + if err != nil { s.logger.Error("Failed to get database stats", "error", err) writeJSONError(w, http.StatusInternalServerError, err.Error()) return - case dbStats = <-statsChan: - // Success } uptime := time.Since(metrics.ConnectedSince).Truncate(time.Second).String() @@ -219,13 +196,6 @@ func (s *Server) handleStatusJSON() http.HandlerFunc { const bitsPerMegabit = 1000000.0 - // Get route counts from database - ipv4Routes, ipv6Routes, err := s.db.GetLiveRouteCountsContext(ctx) - if err != nil { - s.logger.Warn("Failed to get live route counts", "error", err) - // Continue with zero counts - } - // Get route update metrics routeMetrics := s.streamer.GetMetricsTracker().GetRouteMetrics() @@ -259,8 +229,8 @@ func (s *Server) handleStatusJSON() http.HandlerFunc { Peers: dbStats.Peers, DatabaseSizeBytes: dbStats.FileSizeBytes, LiveRoutes: dbStats.LiveRoutes, - IPv4Routes: ipv4Routes, - IPv6Routes: ipv6Routes, + IPv4Routes: dbStats.IPv4Routes, + IPv6Routes: dbStats.IPv6Routes, OldestRoute: dbStats.OldestRoute, NewestRoute: dbStats.NewestRoute, IPv4UpdatesPerSec: routeMetrics.IPv4UpdatesPerSec, @@ -400,36 +370,14 @@ func (s *Server) handleStats() http.HandlerFunc { metrics := s.streamer.GetMetrics() - // Get database stats with timeout. The channels are buffered so the - // goroutine's send never blocks if the timeout wins and nothing here - // receives; otherwise it would block forever and leak. - statsChan := make(chan database.Stats, 1) - errChan := make(chan error, 1) - - go func() { - dbStats, err := s.db.GetStatsContext(ctx) - if err != nil { - s.logger.Debug("Database stats query failed", "error", err) - errChan <- err - - return - } - statsChan <- dbStats - }() - - var dbStats database.Stats - select { - case <-ctx.Done(): - s.logger.Error("Database stats timeout") - // Don't write response here - timeout middleware already handles it - return - case err := <-errChan: + // Serve database statistics from the cache, which runs the table scans at + // most once per interval so this request does not. + dbStats, err := s.stats.get() + if err != nil { s.logger.Error("Failed to get database stats", "error", err) writeJSONError(w, http.StatusInternalServerError, err.Error()) return - case dbStats = <-statsChan: - // Success } uptime := time.Since(metrics.ConnectedSince).Truncate(time.Second).String() @@ -439,13 +387,6 @@ func (s *Server) handleStats() http.HandlerFunc { const bitsPerMegabit = 1000000.0 - // Get route counts from database - ipv4Routes, ipv6Routes, err := s.db.GetLiveRouteCountsContext(ctx) - if err != nil { - s.logger.Warn("Failed to get live route counts", "error", err) - // Continue with zero counts - } - // Get route update metrics routeMetrics := s.streamer.GetMetricsTracker().GetRouteMetrics() @@ -537,8 +478,8 @@ func (s *Server) handleStats() http.HandlerFunc { Peers: dbStats.Peers, DatabaseSizeBytes: dbStats.FileSizeBytes, LiveRoutes: dbStats.LiveRoutes, - IPv4Routes: ipv4Routes, - IPv6Routes: ipv6Routes, + IPv4Routes: dbStats.IPv4Routes, + IPv6Routes: dbStats.IPv6Routes, OldestRoute: dbStats.OldestRoute, NewestRoute: dbStats.NewestRoute, IPv4UpdatesPerSec: routeMetrics.IPv4UpdatesPerSec, diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index 2ac0f43..ebecaac 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -4,9 +4,8 @@ import ( "context" "net/http" "net/http/httptest" - "runtime" + "sync/atomic" "testing" - "time" "git.eeqj.de/sneak/routewatch/internal/database" "git.eeqj.de/sneak/routewatch/internal/logger" @@ -14,86 +13,46 @@ import ( "git.eeqj.de/sneak/routewatch/internal/streamer" ) -// blockingStatsDB embeds database.Store (left nil) and overrides only -// GetStatsContext, which blocks until release is closed. The stats handlers -// call it in a goroutine; every other Store method is unused on the timeout -// path and would panic if called. -type blockingStatsDB struct { +// countingStatsDB embeds database.Store (left nil) and overrides only +// GetStatsContext, counting how many times it runs. The stats handlers read +// their database statistics through the cache, which calls this; every other +// Store method is unused on the stats path and would panic if called. +type countingStatsDB struct { database.Store - release chan struct{} + calls *atomic.Int64 } -func (d blockingStatsDB) GetStatsContext(_ context.Context) (database.Stats, error) { - <-d.release +func (d countingStatsDB) GetStatsContext(_ context.Context) (database.Stats, error) { + d.calls.Add(1) return database.Stats{}, nil } -// TestStatsHandlersDoNotLeakOnTimeout drives each stats handler repeatedly with -// a request whose context times out before the database responds, then releases -// the blocked queries and asserts the goroutine count returns to its starting -// value. Before the fix the per-request goroutine sent on an unbuffered channel -// that nothing received once the timeout won, so it blocked forever and every -// poll leaked one goroutine. -func TestStatsHandlersDoNotLeakOnTimeout(t *testing.T) { - release := make(chan struct{}) - db := blockingStatsDB{release: release} +// TestStatsHandlersServeFromCache drives both stats handlers many times and +// checks that they answer 200 while the database statistics are computed at most +// once within the refresh interval. Before the fix each request ran the counts +// and MIN/MAX scans itself, which took the full timeout and returned 500 once +// the database grew large. +func TestStatsHandlersServeFromCache(t *testing.T) { + var calls atomic.Int64 + db := countingStatsDB{calls: &calls} s := New(db, streamer.New(logger.New(), metrics.New()), logger.New()) - handlers := map[string]http.HandlerFunc{ - "status.json": s.handleStatusJSON(), - "stats": s.handleStats(), - } + handlers := []http.HandlerFunc{s.handleStatusJSON(), s.handleStats()} - baseline := settledGoroutineCount() - - const ( - iterations = 20 - requestTimeout = 50 * time.Millisecond - ) + const iterations = 20 for _, handler := range handlers { for range iterations { - ctx, cancel := context.WithTimeout(context.Background(), requestTimeout) - req := httptest.NewRequest(http.MethodGet, "/", nil).WithContext(ctx) - handler(httptest.NewRecorder(), req) - cancel() + req := httptest.NewRequest(http.MethodGet, "/", nil) + rec := httptest.NewRecorder() + handler(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("handler returned %d, want %d", rec.Code, http.StatusOK) + } } } - // Let the blocked queries finish; with buffered channels each goroutine's - // send now succeeds and the goroutine exits. - close(release) - - if !waitForGoroutines(baseline) { - t.Fatalf("goroutines did not return to baseline %d, got %d", - baseline, runtime.NumGoroutine()) + if got := calls.Load(); got != 1 { + t.Fatalf("GetStatsContext ran %d times, want 1 within the interval", got) } } - -// settledGoroutineCount lets transient goroutines finish, then reports the -// current count. -func settledGoroutineCount() int { - prev := runtime.NumGoroutine() - for range 20 { - time.Sleep(10 * time.Millisecond) - cur := runtime.NumGoroutine() - if cur == prev { - return cur - } - prev = cur - } - - return prev -} - -// waitForGoroutines waits until the goroutine count drops to target or below. -func waitForGoroutines(target int) bool { - for range 100 { - if runtime.NumGoroutine() <= target { - return true - } - time.Sleep(10 * time.Millisecond) - } - - return false -} diff --git a/internal/server/server.go b/internal/server/server.go index afdd426..84430b7 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -35,6 +35,7 @@ type Server struct { logger *logger.Logger srv *http.Server asnFetcher ASNFetcher + stats *statsCache } // New creates a new HTTP server @@ -44,6 +45,9 @@ func New(db database.Store, streamer *streamer.Streamer, logger *logger.Logger) streamer: streamer, logger: logger, } + s.stats = newStatsCache(func(ctx context.Context) (database.Stats, error) { + return s.db.GetStatsContext(ctx) + }) s.setupRoutes() diff --git a/internal/server/statscache.go b/internal/server/statscache.go new file mode 100644 index 0000000..7a08a80 --- /dev/null +++ b/internal/server/statscache.go @@ -0,0 +1,112 @@ +package server + +import ( + "context" + "sync" + "time" + + "git.eeqj.de/sneak/routewatch/internal/database" +) + +const ( + // statsRefreshInterval is how often the cached database statistics are + // recomputed. The scans behind GetStatsContext grow with the tables, so a + // request serves the cached copy instead of running them. + statsRefreshInterval = 30 * time.Second + + // statsComputeTimeout bounds a single statistics computation so a stuck scan + // cannot block the refresh forever. + statsComputeTimeout = 20 * time.Second +) + +// statsFetch computes fresh statistics. It is the expensive database scan that +// the cache runs at most once per interval. +type statsFetch func(ctx context.Context) (database.Stats, error) + +// statsCache serves the most recent database statistics and recomputes them at +// most once per interval. The first request computes synchronously so it has +// real data to return; afterwards requests serve the cached copy immediately +// and a stale copy triggers a single background refresh, so no request waits on +// the scans. +type statsCache struct { + fetch statsFetch + interval time.Duration + now func() time.Time + + mu sync.Mutex + stats database.Stats + haveStats bool + fetchedAt time.Time + refreshing bool +} + +// newStatsCache returns a cache that recomputes statistics with fetch no more +// than once per statsRefreshInterval. +func newStatsCache(fetch statsFetch) *statsCache { + return &statsCache{ + fetch: fetch, + interval: statsRefreshInterval, + now: time.Now, + } +} + +// get returns the cached statistics. On the first call it computes them +// synchronously and returns any error. Later calls return the cached copy, and +// when that copy is older than the interval they start one background refresh. +func (c *statsCache) get() (database.Stats, error) { + c.mu.Lock() + + if !c.haveStats { + // Cold start: compute once under the lock so concurrent first callers + // wait for this single computation rather than each starting their own. + stats, err := c.compute() + if err != nil { + c.mu.Unlock() + + return database.Stats{}, err + } + c.store(stats) + c.mu.Unlock() + + return stats, nil + } + + if c.now().Sub(c.fetchedAt) >= c.interval && !c.refreshing { + c.refreshing = true + go c.refresh() + } + + stats := c.stats + c.mu.Unlock() + + return stats, nil +} + +// refresh recomputes the statistics in the background and replaces the cached +// copy. A failed computation leaves the previous copy in place. +func (c *statsCache) refresh() { + stats, err := c.compute() + + c.mu.Lock() + defer c.mu.Unlock() + c.refreshing = false + if err == nil { + c.store(stats) + } +} + +// compute runs the fetch with its own bounded context, independent of any +// request, so one request's cancellation cannot abort a shared refresh. +func (c *statsCache) compute() (database.Stats, error) { + ctx, cancel := context.WithTimeout(context.Background(), statsComputeTimeout) + defer cancel() + + return c.fetch(ctx) +} + +// store records a fresh result. The caller must hold the mutex. +func (c *statsCache) store(stats database.Stats) { + c.stats = stats + c.haveStats = true + c.fetchedAt = c.now() +} diff --git a/internal/server/statscache_test.go b/internal/server/statscache_test.go new file mode 100644 index 0000000..75a96e2 --- /dev/null +++ b/internal/server/statscache_test.go @@ -0,0 +1,173 @@ +package server + +import ( + "context" + "errors" + "runtime" + "sync/atomic" + "testing" + "time" + + "git.eeqj.de/sneak/routewatch/internal/database" +) + +// testClock is a concurrency-safe clock the cache tests advance by hand, so the +// interval boundary is exercised without waiting real time. +type testClock struct { + ns atomic.Int64 +} + +func (c *testClock) now() time.Time { return time.Unix(0, c.ns.Load()) } +func (c *testClock) advance(d time.Duration) { c.ns.Add(int64(d)) } + +// waitForCalls waits until calls reaches want, giving a background refresh time +// to finish. +func waitForCalls(calls *atomic.Int64, want int64) bool { + const attempts = 200 + for range attempts { + if calls.Load() >= want { + return true + } + time.Sleep(5 * time.Millisecond) + } + + return false +} + +// TestStatsCacheComputesOncePerInterval is the core guarantee: many reads in a +// row run the expensive fetch at most once per interval, and crossing the +// interval boundary allows exactly one more computation. +func TestStatsCacheComputesOncePerInterval(t *testing.T) { + clk := &testClock{} + clk.ns.Store(int64(time.Hour)) // start at a non-zero instant + + var calls atomic.Int64 + c := newStatsCache(func(_ context.Context) (database.Stats, error) { + calls.Add(1) + + return database.Stats{}, nil + }) + c.now = clk.now + + const reads = 50 + for range reads { + if _, err := c.get(); err != nil { + t.Fatalf("get returned error: %v", err) + } + } + if got := calls.Load(); got != 1 { + t.Fatalf("fetch ran %d times within the interval, want 1", got) + } + + // Cross the interval: the next read serves the stale copy and starts one + // background refresh. + clk.advance(c.interval) + if _, err := c.get(); err != nil { + t.Fatalf("get after interval returned error: %v", err) + } + if !waitForCalls(&calls, 2) { + t.Fatalf("background refresh did not run, fetch ran %d times", calls.Load()) + } + + for range reads { + if _, err := c.get(); err != nil { + t.Fatalf("get returned error: %v", err) + } + } + if got := calls.Load(); got != 2 { + t.Fatalf("fetch ran %d times across one interval boundary, want 2", got) + } +} + +// TestStatsCacheColdStartReturnsError checks the first computation's error +// reaches the caller, since there is no cached copy to serve instead. +func TestStatsCacheColdStartReturnsError(t *testing.T) { + wantErr := errors.New("boom") + c := newStatsCache(func(_ context.Context) (database.Stats, error) { + return database.Stats{}, wantErr + }) + + if _, err := c.get(); !errors.Is(err, wantErr) { + t.Fatalf("get returned %v, want %v", err, wantErr) + } +} + +// TestStatsCacheServesLastGoodCopyOnRefreshError checks that once a copy exists, +// a later failing refresh does not surface an error or drop the good data. +func TestStatsCacheServesLastGoodCopyOnRefreshError(t *testing.T) { + clk := &testClock{} + clk.ns.Store(int64(time.Hour)) + + const wantASNs = 7 + var calls atomic.Int64 + var failing atomic.Bool + c := newStatsCache(func(_ context.Context) (database.Stats, error) { + calls.Add(1) + if failing.Load() { + return database.Stats{}, errors.New("boom") + } + + return database.Stats{ASNs: wantASNs}, nil + }) + c.now = clk.now + + got, err := c.get() + if err != nil || got.ASNs != wantASNs { + t.Fatalf("cold start returned (%+v, %v), want ASNs=%d, nil", got, err, wantASNs) + } + + failing.Store(true) + clk.advance(c.interval) + + got, err = c.get() + if err != nil { + t.Fatalf("get during failing refresh returned error: %v", err) + } + if got.ASNs != wantASNs { + t.Fatalf("get returned ASNs=%d, want the last good copy %d", got.ASNs, wantASNs) + } + if !waitForCalls(&calls, 2) { + t.Fatalf("refresh was not attempted, fetch ran %d times", calls.Load()) + } +} + +// TestStatsCacheBackgroundRefreshDoesNotLeak forces many stale refreshes and +// checks the goroutine count returns to its starting value. +func TestStatsCacheBackgroundRefreshDoesNotLeak(t *testing.T) { + clk := &testClock{} + clk.ns.Store(int64(time.Hour)) + + var calls atomic.Int64 + c := newStatsCache(func(_ context.Context) (database.Stats, error) { + calls.Add(1) + + return database.Stats{}, nil + }) + c.now = clk.now + + if _, err := c.get(); err != nil { + t.Fatalf("cold start returned error: %v", err) + } + + baseline := runtime.NumGoroutine() + + const rounds = 20 + for i := range rounds { + clk.advance(c.interval) + if _, err := c.get(); err != nil { + t.Fatalf("get returned error: %v", err) + } + if !waitForCalls(&calls, int64(i+2)) { + t.Fatalf("refresh %d did not run", i) + } + } + + const settleAttempts = 100 + for range settleAttempts { + if runtime.NumGoroutine() <= baseline { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("goroutines did not settle to baseline %d, got %d", baseline, runtime.NumGoroutine()) +}