diff --git a/internal/server/handlers.go b/internal/server/handlers.go index 98ff6b2..d04d0fd 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -179,9 +179,11 @@ func (s *Server) handleStatusJSON() http.HandlerFunc { metrics := s.streamer.GetMetrics() - // Get database stats with timeout - statsChan := make(chan database.Stats) - errChan := make(chan error) + // 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) @@ -398,9 +400,11 @@ func (s *Server) handleStats() http.HandlerFunc { metrics := s.streamer.GetMetrics() - // Get database stats with timeout - statsChan := make(chan database.Stats) - errChan := make(chan error) + // 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) diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go new file mode 100644 index 0000000..2ac0f43 --- /dev/null +++ b/internal/server/handlers_test.go @@ -0,0 +1,99 @@ +package server + +import ( + "context" + "net/http" + "net/http/httptest" + "runtime" + "testing" + "time" + + "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" +) + +// 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()) + + 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()) + } +} + +// 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 +} diff --git a/internal/streamer/streamer.go b/internal/streamer/streamer.go index d345a9a..42d6547 100644 --- a/internal/streamer/streamer.go +++ b/internal/streamer/streamer.go @@ -106,6 +106,7 @@ type handlerInfo struct { type Streamer struct { logger *logger.Logger client *http.Client + url string handlers []*handlerInfo rawHandler RawMessageHandler mu sync.RWMutex @@ -124,6 +125,7 @@ type Streamer struct { func New(logger *logger.Logger, metrics *metrics.Tracker) *Streamer { return &Streamer{ logger: logger, + url: risLiveURL, client: &http.Client{ Timeout: 0, // No timeout for streaming Transport: &http.Transport{ @@ -463,7 +465,14 @@ func (s *Streamer) streamWithReconnect(ctx context.Context) { } func (s *Streamer) stream(ctx context.Context) error { - req, err := http.NewRequestWithContext(ctx, "GET", risLiveURL, nil) + // connCtx is scoped to this single connection: cancelling it when stream + // returns stops the ticker goroutines below, so a reconnect does not leak + // them. Without this they would live until the streamer's lifetime context + // is cancelled, leaking two per reconnect. + connCtx, connCancel := context.WithCancel(ctx) + defer connCancel() + + req, err := http.NewRequestWithContext(ctx, "GET", s.url, nil) if err != nil { return fmt.Errorf("failed to create request: %w", err) } @@ -516,7 +525,7 @@ func (s *Streamer) stream(ctx context.Context) error { select { case <-metricsTicker.C: s.logMetrics() - case <-ctx.Done(): + case <-connCtx.Done(): return } } @@ -536,7 +545,7 @@ func (s *Streamer) stream(ctx context.Context) error { s.metrics.RecordWireBytes(delta) lastWireBytes = currentBytes } - case <-ctx.Done(): + case <-connCtx.Done(): return } } diff --git a/internal/streamer/streamer_test.go b/internal/streamer/streamer_test.go index 1c788a4..3f0b4aa 100644 --- a/internal/streamer/streamer_test.go +++ b/internal/streamer/streamer_test.go @@ -1,7 +1,12 @@ package streamer import ( + "context" + "net/http" + "net/http/httptest" + "runtime" "testing" + "time" "git.eeqj.de/sneak/routewatch/internal/logger" "git.eeqj.de/sneak/routewatch/internal/metrics" @@ -32,3 +37,68 @@ func TestNewStreamer(t *testing.T) { t.Error("metrics tracker not set correctly") } } + +// TestStreamDoesNotLeakTickersAcrossReconnects drives many short-lived +// connections (each stream call is one reconnect cycle) and asserts the +// goroutine count returns to its starting value. Each connection starts two +// ticker goroutines; before the fix they lived until the streamer's lifetime +// context was cancelled, so every reconnect leaked two. +func TestStreamDoesNotLeakTickersAcrossReconnects(t *testing.T) { + // The handler returns immediately, so the response body is empty and each + // stream call ends at once, standing in for a dropped connection. + srv := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {})) + defer srv.Close() + + s := New(logger.New(), metrics.New()) + s.url = srv.URL + + // One warm-up connection so any persistent HTTP transport goroutine exists + // before we take the baseline. + if err := s.stream(context.Background()); err != nil { + t.Fatalf("warm-up stream returned error: %v", err) + } + s.client.CloseIdleConnections() + + baseline := settledGoroutineCount() + + const reconnects = 20 + for range reconnects { + if err := s.stream(context.Background()); err != nil { + t.Fatalf("stream returned error: %v", err) + } + } + s.client.CloseIdleConnections() + + if !waitForGoroutines(baseline) { + t.Fatalf("goroutines did not return to baseline %d after %d reconnects, got %d", + baseline, reconnects, 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 +}