Serve /api/v1/stats from a cache and index-scan the route timestamps (closes #27)
check / check (push) Failing after 0s
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
This commit is contained in:
@@ -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")
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
@@ -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,
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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