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()) }