Fix two goroutine leaks: stats handlers on timeout, streamer tickers on reconnect (closes #12) #19
@@ -179,9 +179,11 @@ func (s *Server) handleStatusJSON() http.HandlerFunc {
|
|||||||
|
|
||||||
metrics := s.streamer.GetMetrics()
|
metrics := s.streamer.GetMetrics()
|
||||||
|
|
||||||
// Get database stats with timeout
|
// Get database stats with timeout. The channels are buffered so the
|
||||||
statsChan := make(chan database.Stats)
|
// goroutine's send never blocks if the timeout wins and nothing here
|
||||||
errChan := make(chan error)
|
// receives; otherwise it would block forever and leak.
|
||||||
|
statsChan := make(chan database.Stats, 1)
|
||||||
|
errChan := make(chan error, 1)
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
dbStats, err := s.db.GetStatsContext(ctx)
|
dbStats, err := s.db.GetStatsContext(ctx)
|
||||||
@@ -398,9 +400,11 @@ func (s *Server) handleStats() http.HandlerFunc {
|
|||||||
|
|
||||||
metrics := s.streamer.GetMetrics()
|
metrics := s.streamer.GetMetrics()
|
||||||
|
|
||||||
// Get database stats with timeout
|
// Get database stats with timeout. The channels are buffered so the
|
||||||
statsChan := make(chan database.Stats)
|
// goroutine's send never blocks if the timeout wins and nothing here
|
||||||
errChan := make(chan error)
|
// receives; otherwise it would block forever and leak.
|
||||||
|
statsChan := make(chan database.Stats, 1)
|
||||||
|
errChan := make(chan error, 1)
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
dbStats, err := s.db.GetStatsContext(ctx)
|
dbStats, err := s.db.GetStatsContext(ctx)
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -106,6 +106,7 @@ type handlerInfo struct {
|
|||||||
type Streamer struct {
|
type Streamer struct {
|
||||||
logger *logger.Logger
|
logger *logger.Logger
|
||||||
client *http.Client
|
client *http.Client
|
||||||
|
url string
|
||||||
handlers []*handlerInfo
|
handlers []*handlerInfo
|
||||||
rawHandler RawMessageHandler
|
rawHandler RawMessageHandler
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
@@ -124,6 +125,7 @@ type Streamer struct {
|
|||||||
func New(logger *logger.Logger, metrics *metrics.Tracker) *Streamer {
|
func New(logger *logger.Logger, metrics *metrics.Tracker) *Streamer {
|
||||||
return &Streamer{
|
return &Streamer{
|
||||||
logger: logger,
|
logger: logger,
|
||||||
|
url: risLiveURL,
|
||||||
client: &http.Client{
|
client: &http.Client{
|
||||||
Timeout: 0, // No timeout for streaming
|
Timeout: 0, // No timeout for streaming
|
||||||
Transport: &http.Transport{
|
Transport: &http.Transport{
|
||||||
@@ -463,7 +465,14 @@ func (s *Streamer) streamWithReconnect(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *Streamer) stream(ctx context.Context) error {
|
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 {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to create request: %w", err)
|
return fmt.Errorf("failed to create request: %w", err)
|
||||||
}
|
}
|
||||||
@@ -516,7 +525,7 @@ func (s *Streamer) stream(ctx context.Context) error {
|
|||||||
select {
|
select {
|
||||||
case <-metricsTicker.C:
|
case <-metricsTicker.C:
|
||||||
s.logMetrics()
|
s.logMetrics()
|
||||||
case <-ctx.Done():
|
case <-connCtx.Done():
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -536,7 +545,7 @@ func (s *Streamer) stream(ctx context.Context) error {
|
|||||||
s.metrics.RecordWireBytes(delta)
|
s.metrics.RecordWireBytes(delta)
|
||||||
lastWireBytes = currentBytes
|
lastWireBytes = currentBytes
|
||||||
}
|
}
|
||||||
case <-ctx.Done():
|
case <-connCtx.Done():
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,7 +1,12 @@
|
|||||||
package streamer
|
package streamer
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"runtime"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"git.eeqj.de/sneak/routewatch/internal/logger"
|
"git.eeqj.de/sneak/routewatch/internal/logger"
|
||||||
"git.eeqj.de/sneak/routewatch/internal/metrics"
|
"git.eeqj.de/sneak/routewatch/internal/metrics"
|
||||||
@@ -32,3 +37,68 @@ func TestNewStreamer(t *testing.T) {
|
|||||||
t.Error("metrics tracker not set correctly")
|
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
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user