Author SHA1 Message Date
sneak cd59cb8a8d Serve /api/v1/stats from a cache and index-scan the route timestamps (closes #27)
check / check (push) Failing after 0s
Once the database passed about 4.5 GiB every stats request ran a COUNT(*)
over each table plus a MIN/MAX union scan of both route tables, took the
full timeout and returned HTTP 500, so the status page went blank.

The server now keeps the last database statistics in memory and recomputes
them at most once every 30 seconds; requests serve the cached copy and a
stale copy triggers a single background refresh, so no request runs the
scans. The route-count split is folded into the cached stats, removing the
separate per-request live-route count query.

The oldest/newest route timestamps now read one row from each end of the
last_updated index instead of scanning both tables, and select the column
directly so the driver parses it into time.Time; the old aggregate returned
an untyped string that failed to scan and logged a warning every call.

Model: opus-4-8
2026-09-21 23:32:11 +00:00
8 changed files with 498 additions and 152 deletions
+68 -13
View File
@@ -980,23 +980,20 @@ func (d *Database) GetStatsContext(ctx context.Context) (Stats, error) {
if err != nil { if err != nil {
return stats, fmt.Errorf("failed to count IPv6 routes: %w", err) return stats, fmt.Errorf("failed to count IPv6 routes: %w", err)
} }
stats.IPv4Routes = v4Count
stats.IPv6Routes = v6Count
stats.LiveRoutes = v4Count + v6Count stats.LiveRoutes = v4Count + v6Count
// Get oldest and newest route timestamps // Get oldest and newest route timestamps. Each query reads a single row from
routeTimestampQuery := ` // one end of the last_updated index, so the cost is a log-time index lookup
SELECT MIN(last_updated), MAX(last_updated) FROM ( // rather than a full scan of both route tables. Selecting the last_updated
SELECT last_updated FROM live_routes_v4 // column directly (rather than MIN/MAX, whose result has no column type) lets
UNION ALL // the driver parse the DATETIME value into time.Time; the union scan aggregate
SELECT last_updated FROM live_routes_v6 // used before returned an untyped string and logged a warning on every call.
) stats.OldestRoute, stats.NewestRoute, err = d.routeTimestampRange(ctx)
`
var oldestRoute, newestRoute *time.Time
err = d.db.QueryRowContext(ctx, routeTimestampQuery).Scan(&oldestRoute, &newestRoute)
if err != nil { if err != nil {
// Display-only fields; log but keep the rest of the stats.
d.logger.Warn("Failed to get route timestamps", "error", err) d.logger.Warn("Failed to get route timestamps", "error", err)
} else {
stats.OldestRoute = oldestRoute
stats.NewestRoute = newestRoute
} }
// Get prefix distribution // Get prefix distribution
@@ -1009,6 +1006,64 @@ func (d *Database) GetStatsContext(ctx context.Context) (Stats, error) {
return stats, nil 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 // UpsertLiveRoute inserts or updates a live route
func (d *Database) UpsertLiveRoute(route *LiveRoute) error { func (d *Database) UpsertLiveRoute(route *LiveRoute) error {
d.lock("UpsertLiveRoute") d.lock("UpsertLiveRoute")
+100
View File
@@ -10,6 +10,7 @@ import (
"git.eeqj.de/sneak/routewatch/internal/config" "git.eeqj.de/sneak/routewatch/internal/config"
"git.eeqj.de/sneak/routewatch/internal/logger" "git.eeqj.de/sneak/routewatch/internal/logger"
"github.com/google/uuid"
) )
// tempStoreMemory is the PRAGMA temp_store value meaning "hold temp B-trees in // tempStoreMemory is the PRAGMA temp_store value meaning "hold temp B-trees in
@@ -425,6 +426,105 @@ func TestBatchWriteDuringCheckpoint(t *testing.T) {
wg.Wait() 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) { func BenchmarkIPToUint32(b *testing.B) {
ip := net.ParseIP("192.168.1.1") ip := net.ParseIP("192.168.1.1")
b.ResetTimer() b.ResetTimer()
+2
View File
@@ -18,6 +18,8 @@ type Stats struct {
Peers int Peers int
FileSizeBytes int64 FileSizeBytes int64
LiveRoutes int LiveRoutes int
IPv4Routes int
IPv6Routes int
OldestRoute *time.Time OldestRoute *time.Time
NewestRoute *time.Time NewestRoute *time.Time
IPv4PrefixDistribution []PrefixDistribution IPv4PrefixDistribution []PrefixDistribution
+12 -71
View File
@@ -179,37 +179,14 @@ func (s *Server) handleStatusJSON() http.HandlerFunc {
metrics := s.streamer.GetMetrics() metrics := s.streamer.GetMetrics()
// Get database stats with timeout. The channels are buffered so the // Serve database statistics from the cache, which runs the table scans at
// goroutine's send never blocks if the timeout wins and nothing here // most once per interval so this request does not.
// receives; otherwise it would block forever and leak. dbStats, err := s.stats.get()
statsChan := make(chan database.Stats, 1) if err != nil {
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:
s.logger.Error("Failed to get database stats", "error", err) s.logger.Error("Failed to get database stats", "error", err)
writeJSONError(w, http.StatusInternalServerError, err.Error()) writeJSONError(w, http.StatusInternalServerError, err.Error())
return return
case dbStats = <-statsChan:
// Success
} }
uptime := time.Since(metrics.ConnectedSince).Truncate(time.Second).String() uptime := time.Since(metrics.ConnectedSince).Truncate(time.Second).String()
@@ -219,13 +196,6 @@ func (s *Server) handleStatusJSON() http.HandlerFunc {
const bitsPerMegabit = 1000000.0 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 // Get route update metrics
routeMetrics := s.streamer.GetMetricsTracker().GetRouteMetrics() routeMetrics := s.streamer.GetMetricsTracker().GetRouteMetrics()
@@ -259,8 +229,8 @@ func (s *Server) handleStatusJSON() http.HandlerFunc {
Peers: dbStats.Peers, Peers: dbStats.Peers,
DatabaseSizeBytes: dbStats.FileSizeBytes, DatabaseSizeBytes: dbStats.FileSizeBytes,
LiveRoutes: dbStats.LiveRoutes, LiveRoutes: dbStats.LiveRoutes,
IPv4Routes: ipv4Routes, IPv4Routes: dbStats.IPv4Routes,
IPv6Routes: ipv6Routes, IPv6Routes: dbStats.IPv6Routes,
OldestRoute: dbStats.OldestRoute, OldestRoute: dbStats.OldestRoute,
NewestRoute: dbStats.NewestRoute, NewestRoute: dbStats.NewestRoute,
IPv4UpdatesPerSec: routeMetrics.IPv4UpdatesPerSec, IPv4UpdatesPerSec: routeMetrics.IPv4UpdatesPerSec,
@@ -400,36 +370,14 @@ func (s *Server) handleStats() http.HandlerFunc {
metrics := s.streamer.GetMetrics() metrics := s.streamer.GetMetrics()
// Get database stats with timeout. The channels are buffered so the // Serve database statistics from the cache, which runs the table scans at
// goroutine's send never blocks if the timeout wins and nothing here // most once per interval so this request does not.
// receives; otherwise it would block forever and leak. dbStats, err := s.stats.get()
statsChan := make(chan database.Stats, 1) if err != nil {
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:
s.logger.Error("Failed to get database stats", "error", err) s.logger.Error("Failed to get database stats", "error", err)
writeJSONError(w, http.StatusInternalServerError, err.Error()) writeJSONError(w, http.StatusInternalServerError, err.Error())
return return
case dbStats = <-statsChan:
// Success
} }
uptime := time.Since(metrics.ConnectedSince).Truncate(time.Second).String() uptime := time.Since(metrics.ConnectedSince).Truncate(time.Second).String()
@@ -439,13 +387,6 @@ func (s *Server) handleStats() http.HandlerFunc {
const bitsPerMegabit = 1000000.0 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 // Get route update metrics
routeMetrics := s.streamer.GetMetricsTracker().GetRouteMetrics() routeMetrics := s.streamer.GetMetricsTracker().GetRouteMetrics()
@@ -537,8 +478,8 @@ func (s *Server) handleStats() http.HandlerFunc {
Peers: dbStats.Peers, Peers: dbStats.Peers,
DatabaseSizeBytes: dbStats.FileSizeBytes, DatabaseSizeBytes: dbStats.FileSizeBytes,
LiveRoutes: dbStats.LiveRoutes, LiveRoutes: dbStats.LiveRoutes,
IPv4Routes: ipv4Routes, IPv4Routes: dbStats.IPv4Routes,
IPv6Routes: ipv6Routes, IPv6Routes: dbStats.IPv6Routes,
OldestRoute: dbStats.OldestRoute, OldestRoute: dbStats.OldestRoute,
NewestRoute: dbStats.NewestRoute, NewestRoute: dbStats.NewestRoute,
IPv4UpdatesPerSec: routeMetrics.IPv4UpdatesPerSec, IPv4UpdatesPerSec: routeMetrics.IPv4UpdatesPerSec,
+27 -68
View File
@@ -4,9 +4,8 @@ import (
"context" "context"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"runtime" "sync/atomic"
"testing" "testing"
"time"
"git.eeqj.de/sneak/routewatch/internal/database" "git.eeqj.de/sneak/routewatch/internal/database"
"git.eeqj.de/sneak/routewatch/internal/logger" "git.eeqj.de/sneak/routewatch/internal/logger"
@@ -14,86 +13,46 @@ import (
"git.eeqj.de/sneak/routewatch/internal/streamer" "git.eeqj.de/sneak/routewatch/internal/streamer"
) )
// blockingStatsDB embeds database.Store (left nil) and overrides only // countingStatsDB embeds database.Store (left nil) and overrides only
// GetStatsContext, which blocks until release is closed. The stats handlers // GetStatsContext, counting how many times it runs. The stats handlers read
// call it in a goroutine; every other Store method is unused on the timeout // their database statistics through the cache, which calls this; every other
// path and would panic if called. // Store method is unused on the stats path and would panic if called.
type blockingStatsDB struct { type countingStatsDB struct {
database.Store database.Store
release chan struct{} calls *atomic.Int64
} }
func (d blockingStatsDB) GetStatsContext(_ context.Context) (database.Stats, error) { func (d countingStatsDB) GetStatsContext(_ context.Context) (database.Stats, error) {
<-d.release d.calls.Add(1)
return database.Stats{}, nil return database.Stats{}, nil
} }
// TestStatsHandlersDoNotLeakOnTimeout drives each stats handler repeatedly with // TestStatsHandlersServeFromCache drives both stats handlers many times and
// a request whose context times out before the database responds, then releases // checks that they answer 200 while the database statistics are computed at most
// the blocked queries and asserts the goroutine count returns to its starting // once within the refresh interval. Before the fix each request ran the counts
// value. Before the fix the per-request goroutine sent on an unbuffered channel // and MIN/MAX scans itself, which took the full timeout and returned 500 once
// that nothing received once the timeout won, so it blocked forever and every // the database grew large.
// poll leaked one goroutine. func TestStatsHandlersServeFromCache(t *testing.T) {
func TestStatsHandlersDoNotLeakOnTimeout(t *testing.T) { var calls atomic.Int64
release := make(chan struct{}) db := countingStatsDB{calls: &calls}
db := blockingStatsDB{release: release}
s := New(db, streamer.New(logger.New(), metrics.New()), logger.New()) s := New(db, streamer.New(logger.New(), metrics.New()), logger.New())
handlers := map[string]http.HandlerFunc{ handlers := []http.HandlerFunc{s.handleStatusJSON(), s.handleStats()}
"status.json": s.handleStatusJSON(),
"stats": s.handleStats(),
}
baseline := settledGoroutineCount() const iterations = 20
const (
iterations = 20
requestTimeout = 50 * time.Millisecond
)
for _, handler := range handlers { for _, handler := range handlers {
for range iterations { for range iterations {
ctx, cancel := context.WithTimeout(context.Background(), requestTimeout) req := httptest.NewRequest(http.MethodGet, "/", nil)
req := httptest.NewRequest(http.MethodGet, "/", nil).WithContext(ctx) rec := httptest.NewRecorder()
handler(httptest.NewRecorder(), req) handler(rec, req)
cancel() 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 if got := calls.Load(); got != 1 {
// send now succeeds and the goroutine exits. t.Fatalf("GetStatsContext ran %d times, want 1 within the interval", got)
close(release)
if !waitForGoroutines(baseline) {
t.Fatalf("goroutines did not return to baseline %d, got %d",
baseline, runtime.NumGoroutine())
} }
} }
// 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
}
+4
View File
@@ -35,6 +35,7 @@ type Server struct {
logger *logger.Logger logger *logger.Logger
srv *http.Server srv *http.Server
asnFetcher ASNFetcher asnFetcher ASNFetcher
stats *statsCache
} }
// New creates a new HTTP server // New creates a new HTTP server
@@ -44,6 +45,9 @@ func New(db database.Store, streamer *streamer.Streamer, logger *logger.Logger)
streamer: streamer, streamer: streamer,
logger: logger, logger: logger,
} }
s.stats = newStatsCache(func(ctx context.Context) (database.Stats, error) {
return s.db.GetStatsContext(ctx)
})
s.setupRoutes() s.setupRoutes()
+112
View File
@@ -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()
}
+173
View File
@@ -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())
}