Serve /api/v1/stats from a cache and index-scan the route timestamps (closes #27) #28
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
+12
-71
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
Reference in New Issue
Block a user