package server import ( "context" "encoding/json" "net/http" "net/http/httptest" "runtime" "slices" "testing" "time" "git.eeqj.de/sneak/routewatch/internal/config" "git.eeqj.de/sneak/routewatch/internal/database" "git.eeqj.de/sneak/routewatch/internal/logger" "git.eeqj.de/sneak/routewatch/internal/metrics" "git.eeqj.de/sneak/routewatch/internal/streamer" "github.com/google/uuid" ) // 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 { database.Store release chan struct{} } func (d blockingStatsDB) GetStatsContext(_ context.Context) (database.Stats, error) { <-d.release 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} s := New(db, streamer.New(logger.New(), metrics.New()), logger.New(), &config.Config{}) handlers := map[string]http.HandlerFunc{ "status.json": s.handleStatusJSON(), "stats": s.handleStats(), } baseline := settledGoroutineCount() const ( iterations = 20 requestTimeout = 50 * time.Millisecond ) 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() } } // 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()) } } // TestStatsHandlersAnswerFromMemory checks that both stats handlers answer 200 // with the live route counts and the prefix distribution while the database is // closed, so that any query would fail: the request path reads them from // memory. The prefix distribution query it used to run read every live route // and, on a large database, took the whole 4-second deadline, so // /api/v1/stats answered 500 (https://git.eeqj.de/sneak/routewatch/issues/30). // The oldest and newest route times still come from one-row lookups at the ends // of an index; with the database closed they are left out of the answer. func TestStatsHandlersAnswerFromMemory(t *testing.T) { db, err := database.New(&config.Config{StateDir: t.TempDir()}, logger.New()) if err != nil { t.Fatalf("database.New: %v", err) } ts := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) start, end, err := database.CalculateIPv4Range("198.51.100.0/24") if err != nil { t.Fatalf("CalculateIPv4Range: %v", err) } if err := db.UpsertLiveRouteBatch([]*database.LiveRoute{ { ID: uuid.New(), Prefix: "198.51.100.0/24", MaskLength: 24, IPVersion: 4, OriginASN: 64500, PeerIP: "192.0.2.1", ASPath: []int{64500}, NextHop: "192.0.2.1", LastUpdated: ts, V4IPStart: &start, V4IPEnd: &end, }, { ID: uuid.New(), Prefix: "2001:db8::/32", MaskLength: 32, IPVersion: 6, OriginASN: 64501, PeerIP: "2001:db8::1", ASPath: []int{64501}, NextHop: "2001:db8::1", LastUpdated: ts, }, }); err != nil { t.Fatalf("UpsertLiveRouteBatch: %v", err) } if err := db.Close(); err != nil { t.Fatalf("Close: %v", err) } s := New(db, streamer.New(logger.New(), metrics.New()), logger.New(), &config.Config{}) handlers := map[string]http.HandlerFunc{ "status.json": s.handleStatusJSON(), "stats": s.handleStats(), } for name, handler := range handlers { rec := httptest.NewRecorder() handler(rec, httptest.NewRequest(http.MethodGet, "/", nil)) if rec.Code != http.StatusOK { t.Errorf("%s: status %d, want %d; body %s", name, rec.Code, http.StatusOK, rec.Body) continue } var body struct { Data struct { IPv4Routes int `json:"ipv4_routes"` IPv6Routes int `json:"ipv6_routes"` IPv4PrefixDistribution []database.PrefixDistribution `json:"ipv4_prefix_distribution"` IPv6PrefixDistribution []database.PrefixDistribution `json:"ipv6_prefix_distribution"` } `json:"data"` } if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { t.Fatalf("%s: decoding the answer: %v", name, err) } if body.Data.IPv4Routes != 1 || body.Data.IPv6Routes != 1 { t.Errorf("%s: routes = (v4 %d, v6 %d), want (1, 1)", name, body.Data.IPv4Routes, body.Data.IPv6Routes) } wantV4 := []database.PrefixDistribution{{MaskLength: 24, Count: 1}} if !slices.Equal(body.Data.IPv4PrefixDistribution, wantV4) { t.Errorf("%s: IPv4 distribution = %v, want %v", name, body.Data.IPv4PrefixDistribution, wantV4) } wantV6 := []database.PrefixDistribution{{MaskLength: 32, Count: 1}} if !slices.Equal(body.Data.IPv6PrefixDistribution, wantV6) { t.Errorf("%s: IPv6 distribution = %v, want %v", name, body.Data.IPv6PrefixDistribution, wantV6) } } } // 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 }