diff options
| author | Jeff Halter <868228+jhalter@users.noreply.github.com> | 2026-05-28 16:22:24 -0700 |
|---|---|---|
| committer | Jeff Halter <868228+jhalter@users.noreply.github.com> | 2026-05-28 16:22:24 -0700 |
| commit | b772019454ebb804c313e717ae98eb68430e780e (patch) | |
| tree | a12ffc1e5de4d7e1612edaea03174043a6c55784 | |
| parent | 2f46f87177234070044b5304ca3e0db471699ad8 (diff) | |
Refactor Stats to typed keys and fix peak-tracking race
Replace the map-backed Stats counter with a fixed [numStats]int array
indexed by a new StatKey enum, eliminating the hand-maintained map
initialization and the parallel string-keyed Values() map.
Fix a check-then-act race in the connection peak tracking: the old
Get-then-Set across two lock acquisitions could let concurrent
connections clobber each other's update. The new atomic Max method does
the compare-and-set under a single lock.
Values() now returns a typed StatValues struct whose JSON tags preserve
the existing /api/v1/stats wire format.
| -rw-r--r-- | hotline/server.go | 6 | ||||
| -rw-r--r-- | hotline/server_test.go | 6 | ||||
| -rw-r--r-- | hotline/stats.go | 98 | ||||
| -rw-r--r-- | hotline/stats_test.go | 125 | ||||
| -rw-r--r-- | internal/mobius/api_test.go | 33 |
5 files changed, 164 insertions, 104 deletions
diff --git a/hotline/server.go b/hotline/server.go index dcb90e7..d757062 100644 --- a/hotline/server.go +++ b/hotline/server.go @@ -157,7 +157,7 @@ func NewServer(options ...Option) (*Server, error) { return &server, nil } -func (s *Server) CurrentStats() map[string]interface{} { +func (s *Server) CurrentStats() StatValues { return s.Stats.Values() } @@ -648,9 +648,7 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser c.Server.Stats.Increment(StatConnectionCounter, StatCurrentlyConnected) defer c.Server.Stats.Decrement(StatCurrentlyConnected) - if len(s.ClientMgr.List()) > c.Server.Stats.Get(StatConnectionPeak) { - c.Server.Stats.Set(StatConnectionPeak, len(s.ClientMgr.List())) - } + c.Server.Stats.Max(StatConnectionPeak, len(s.ClientMgr.List())) // Scan for new transactions and handle them as they come in. for scanner.Scan() { diff --git a/hotline/server_test.go b/hotline/server_test.go index 4bb3669..bebde8f 100644 --- a/hotline/server_test.go +++ b/hotline/server_test.go @@ -759,9 +759,9 @@ func TestServer_CurrentStats(t *testing.T) { srv := &Server{Stats: stats} result := srv.CurrentStats() - assert.Equal(t, 1, result["CurrentlyConnected"]) - assert.Equal(t, 2, result["DownloadCounter"]) - assert.Equal(t, 0, result["UploadsInProgress"]) + assert.Equal(t, 1, result.CurrentlyConnected) + assert.Equal(t, 2, result.DownloadCounter) + assert.Equal(t, 0, result.UploadsInProgress) } func TestServer_sendTransaction(t *testing.T) { diff --git a/hotline/stats.go b/hotline/stats.go index 9731601..f272f5b 100644 --- a/hotline/stats.go +++ b/hotline/stats.go @@ -5,9 +5,12 @@ import ( "time" ) -// Stat counter keys +// StatKey identifies a single stat counter. +type StatKey int + +// Stat counter keys. numStats must remain last; it sizes the counter array. const ( - StatCurrentlyConnected = iota + StatCurrentlyConnected StatKey = iota StatDownloadsInProgress StatUploadsInProgress StatWaitingDownloads @@ -15,40 +18,45 @@ const ( StatConnectionCounter StatDownloadCounter StatUploadCounter + + numStats ) type Counter interface { - Increment(keys ...int) - Decrement(key int) - Set(key, val int) - Get(key int) int - Values() map[string]interface{} + Increment(keys ...StatKey) + Decrement(keys ...StatKey) + Set(key StatKey, val int) + Max(key StatKey, val int) + Get(key StatKey) int + Values() StatValues +} + +// StatValues is a point-in-time snapshot of all counters. Its JSON tags define +// the wire format served at GET /api/v1/stats. +type StatValues struct { + CurrentlyConnected int `json:"CurrentlyConnected"` + DownloadsInProgress int `json:"DownloadsInProgress"` + UploadsInProgress int `json:"UploadsInProgress"` + WaitingDownloads int `json:"WaitingDownloads"` + ConnectionPeak int `json:"ConnectionPeak"` + ConnectionCounter int `json:"ConnectionCounter"` + DownloadCounter int `json:"DownloadCounter"` + UploadCounter int `json:"UploadCounter"` + Since time.Time `json:"Since"` } type Stats struct { - stats map[int]int + stats [numStats]int since time.Time mu sync.RWMutex } func NewStats() *Stats { - return &Stats{ - since: time.Now(), - stats: map[int]int{ - StatCurrentlyConnected: 0, - StatDownloadsInProgress: 0, - StatUploadsInProgress: 0, - StatWaitingDownloads: 0, - StatConnectionPeak: 0, - StatDownloadCounter: 0, - StatUploadCounter: 0, - StatConnectionCounter: 0, - }, - } + return &Stats{since: time.Now()} } -func (s *Stats) Increment(keys ...int) { +func (s *Stats) Increment(keys ...StatKey) { s.mu.Lock() defer s.mu.Unlock() @@ -57,42 +65,56 @@ func (s *Stats) Increment(keys ...int) { } } -func (s *Stats) Decrement(key int) { +func (s *Stats) Decrement(keys ...StatKey) { s.mu.Lock() defer s.mu.Unlock() - if s.stats[key] > 0 { - s.stats[key]-- + for _, key := range keys { + if s.stats[key] > 0 { + s.stats[key]-- + } } } -func (s *Stats) Set(key, val int) { +func (s *Stats) Set(key StatKey, val int) { s.mu.Lock() defer s.mu.Unlock() s.stats[key] = val } -func (s *Stats) Get(key int) int { +// Max sets key to val only if val is greater than the current value. The +// compare-and-set happens under a single lock so concurrent callers cannot +// clobber each other's update. +func (s *Stats) Max(key StatKey, val int) { + s.mu.Lock() + defer s.mu.Unlock() + + if val > s.stats[key] { + s.stats[key] = val + } +} + +func (s *Stats) Get(key StatKey) int { s.mu.RLock() defer s.mu.RUnlock() return s.stats[key] } -func (s *Stats) Values() map[string]interface{} { +func (s *Stats) Values() StatValues { s.mu.RLock() defer s.mu.RUnlock() - return map[string]interface{}{ - "CurrentlyConnected": s.stats[StatCurrentlyConnected], - "DownloadsInProgress": s.stats[StatDownloadsInProgress], - "UploadsInProgress": s.stats[StatUploadsInProgress], - "WaitingDownloads": s.stats[StatWaitingDownloads], - "ConnectionPeak": s.stats[StatConnectionPeak], - "ConnectionCounter": s.stats[StatConnectionCounter], - "DownloadCounter": s.stats[StatDownloadCounter], - "UploadCounter": s.stats[StatUploadCounter], - "Since": s.since, + return StatValues{ + CurrentlyConnected: s.stats[StatCurrentlyConnected], + DownloadsInProgress: s.stats[StatDownloadsInProgress], + UploadsInProgress: s.stats[StatUploadsInProgress], + WaitingDownloads: s.stats[StatWaitingDownloads], + ConnectionPeak: s.stats[StatConnectionPeak], + ConnectionCounter: s.stats[StatConnectionCounter], + DownloadCounter: s.stats[StatDownloadCounter], + UploadCounter: s.stats[StatUploadCounter], + Since: s.since, } } diff --git a/hotline/stats_test.go b/hotline/stats_test.go index d853dee..38a23c4 100644 --- a/hotline/stats_test.go +++ b/hotline/stats_test.go @@ -1,8 +1,9 @@ package hotline import ( + "encoding/json" + "sync" "testing" - "time" "github.com/stretchr/testify/assert" ) @@ -10,20 +11,20 @@ import ( func TestStats_Increment(t *testing.T) { tests := []struct { name string - keys []int - expected map[int]int + keys []StatKey + expected map[StatKey]int }{ { name: "single key increment", - keys: []int{StatCurrentlyConnected}, - expected: map[int]int{ + keys: []StatKey{StatCurrentlyConnected}, + expected: map[StatKey]int{ StatCurrentlyConnected: 1, }, }, { name: "multiple keys increment", - keys: []int{StatCurrentlyConnected, StatDownloadCounter, StatUploadCounter}, - expected: map[int]int{ + keys: []StatKey{StatCurrentlyConnected, StatDownloadCounter, StatUploadCounter}, + expected: map[StatKey]int{ StatCurrentlyConnected: 1, StatDownloadCounter: 1, StatUploadCounter: 1, @@ -31,8 +32,8 @@ func TestStats_Increment(t *testing.T) { }, { name: "duplicate keys increment", - keys: []int{StatCurrentlyConnected, StatCurrentlyConnected}, - expected: map[int]int{ + keys: []StatKey{StatCurrentlyConnected, StatCurrentlyConnected}, + expected: map[StatKey]int{ StatCurrentlyConnected: 2, }, }, @@ -68,7 +69,7 @@ func TestStats_Decrement(t *testing.T) { tests := []struct { name string setupValue int - key int + key StatKey expected int }{ { @@ -119,7 +120,7 @@ func TestStats_Decrement_Multiple_Calls(t *testing.T) { func TestStats_Set(t *testing.T) { tests := []struct { name string - key int + key StatKey value int expected int }{ @@ -167,7 +168,7 @@ func TestStats_Set(t *testing.T) { func TestStats_Get(t *testing.T) { tests := []struct { name string - key int + key StatKey setValue int expected int }{ @@ -210,7 +211,7 @@ func TestStats_Get(t *testing.T) { func TestStats_Get_Default_Values(t *testing.T) { stats := NewStats() - expectedDefaults := map[int]int{ + expectedDefaults := map[StatKey]int{ StatCurrentlyConnected: 0, StatDownloadsInProgress: 0, StatUploadsInProgress: 0, @@ -232,19 +233,15 @@ func TestStats_Values(t *testing.T) { // Test default values values := stats.Values() - assert.Equal(t, 0, values["CurrentlyConnected"]) - assert.Equal(t, 0, values["DownloadsInProgress"]) - assert.Equal(t, 0, values["UploadsInProgress"]) - assert.Equal(t, 0, values["WaitingDownloads"]) - assert.Equal(t, 0, values["ConnectionPeak"]) - assert.Equal(t, 0, values["ConnectionCounter"]) - assert.Equal(t, 0, values["DownloadCounter"]) - assert.Equal(t, 0, values["UploadCounter"]) - assert.NotNil(t, values["Since"]) - - // Verify Since is a time.Time - _, ok := values["Since"].(time.Time) - assert.True(t, ok, "Since should be a time.Time") + assert.Equal(t, 0, values.CurrentlyConnected) + assert.Equal(t, 0, values.DownloadsInProgress) + assert.Equal(t, 0, values.UploadsInProgress) + assert.Equal(t, 0, values.WaitingDownloads) + assert.Equal(t, 0, values.ConnectionPeak) + assert.Equal(t, 0, values.ConnectionCounter) + assert.Equal(t, 0, values.DownloadCounter) + assert.Equal(t, 0, values.UploadCounter) + assert.False(t, values.Since.IsZero(), "Since should be set") } func TestStats_Values_WithModifiedStats(t *testing.T) { @@ -258,19 +255,23 @@ func TestStats_Values_WithModifiedStats(t *testing.T) { values := stats.Values() - assert.Equal(t, 10, values["CurrentlyConnected"]) - assert.Equal(t, 5, values["DownloadsInProgress"]) - assert.Equal(t, 0, values["UploadsInProgress"]) - assert.Equal(t, 0, values["WaitingDownloads"]) - assert.Equal(t, 0, values["ConnectionPeak"]) - assert.Equal(t, 1, values["ConnectionCounter"]) - assert.Equal(t, 1, values["DownloadCounter"]) - assert.Equal(t, 1, values["UploadCounter"]) + assert.Equal(t, 10, values.CurrentlyConnected) + assert.Equal(t, 5, values.DownloadsInProgress) + assert.Equal(t, 0, values.UploadsInProgress) + assert.Equal(t, 0, values.WaitingDownloads) + assert.Equal(t, 0, values.ConnectionPeak) + assert.Equal(t, 1, values.ConnectionCounter) + assert.Equal(t, 1, values.DownloadCounter) + assert.Equal(t, 1, values.UploadCounter) } -func TestStats_Values_ContainsAllKeys(t *testing.T) { - stats := NewStats() - values := stats.Values() +// TestStats_Values_JSONKeys locks the JSON wire format served at /api/v1/stats. +func TestStats_Values_JSONKeys(t *testing.T) { + b, err := json.Marshal(NewStats().Values()) + assert.NoError(t, err) + + var decoded map[string]interface{} + assert.NoError(t, json.Unmarshal(b, &decoded)) expectedKeys := []string{ "CurrentlyConnected", @@ -285,10 +286,52 @@ func TestStats_Values_ContainsAllKeys(t *testing.T) { } for _, key := range expectedKeys { - _, exists := values[key] - assert.True(t, exists, "Key %s should exist in Values() output", key) + _, exists := decoded[key] + assert.True(t, exists, "Key %s should exist in JSON output", key) } - // Should have exactly 9 keys - assert.Equal(t, 9, len(values)) + assert.Equal(t, len(expectedKeys), len(decoded)) +} + +func TestStats_Max(t *testing.T) { + stats := NewStats() + + stats.Max(StatConnectionPeak, 5) + assert.Equal(t, 5, stats.Get(StatConnectionPeak), "raises to a larger value") + + stats.Max(StatConnectionPeak, 3) + assert.Equal(t, 5, stats.Get(StatConnectionPeak), "no-op for a smaller value") + + stats.Max(StatConnectionPeak, 5) + assert.Equal(t, 5, stats.Get(StatConnectionPeak), "no-op for an equal value") + + stats.Max(StatConnectionPeak, 10) + assert.Equal(t, 10, stats.Get(StatConnectionPeak), "raises again") +} + +// TestStats_Concurrent exercises the counter from many goroutines so the race +// detector can catch unsynchronized access (run with -race). +func TestStats_Concurrent(t *testing.T) { + stats := NewStats() + + const goroutines = 50 + var wg sync.WaitGroup + wg.Add(goroutines) + + for i := 0; i < goroutines; i++ { + go func(n int) { + defer wg.Done() + stats.Increment(StatConnectionCounter) + stats.Max(StatConnectionPeak, n) + stats.Set(StatCurrentlyConnected, n) + _ = stats.Get(StatCurrentlyConnected) + _ = stats.Values() + stats.Decrement(StatConnectionCounter) + }(i) + } + + wg.Wait() + + assert.Equal(t, 0, stats.Get(StatConnectionCounter), "increments and decrements balance out") + assert.Equal(t, goroutines-1, stats.Get(StatConnectionPeak), "peak captures the largest value") } diff --git a/internal/mobius/api_test.go b/internal/mobius/api_test.go index d448b1b..ae9916b 100644 --- a/internal/mobius/api_test.go +++ b/internal/mobius/api_test.go @@ -173,19 +173,15 @@ func (m *mockClientMgr) Delete(id hotline.ClientID) { } type mockCounter struct { - vals map[string]interface{} + vals hotline.StatValues } -func (m *mockCounter) Increment(_ ...int) {} -func (m *mockCounter) Decrement(_ int) {} -func (m *mockCounter) Set(_, _ int) {} -func (m *mockCounter) Get(_ int) int { return 0 } -func (m *mockCounter) Values() map[string]interface{} { - if m.vals == nil { - return map[string]interface{}{} - } - return m.vals -} +func (m *mockCounter) Increment(_ ...hotline.StatKey) {} +func (m *mockCounter) Decrement(_ ...hotline.StatKey) {} +func (m *mockCounter) Set(_ hotline.StatKey, _ int) {} +func (m *mockCounter) Max(_ hotline.StatKey, _ int) {} +func (m *mockCounter) Get(_ hotline.StatKey) int { return 0 } +func (m *mockCounter) Values() hotline.StatValues { return m.vals } // --- Test helper --- @@ -651,9 +647,9 @@ func TestShutdownHandler(t *testing.T) { func TestRenderStats(t *testing.T) { t.Run("returns JSON stats", func(t *testing.T) { srv, _, _, counter := newTestAPIServer(t, "") - counter.vals = map[string]interface{}{ - "connections": 42, - "downloads": 10, + counter.vals = hotline.StatValues{ + CurrentlyConnected: 42, + DownloadCounter: 10, } req := httptest.NewRequest(http.MethodGet, "/api/v1/stats", nil) @@ -666,11 +662,11 @@ func TestRenderStats(t *testing.T) { var stats map[string]interface{} err := json.Unmarshal(rr.Body.Bytes(), &stats) require.NoError(t, err) - assert.Equal(t, float64(42), stats["connections"]) - assert.Equal(t, float64(10), stats["downloads"]) + assert.Equal(t, float64(42), stats["CurrentlyConnected"]) + assert.Equal(t, float64(10), stats["DownloadCounter"]) }) - t.Run("returns empty stats when no data", func(t *testing.T) { + t.Run("returns zero-valued stats when no data", func(t *testing.T) { srv, _, _, _ := newTestAPIServer(t, "") req := httptest.NewRequest(http.MethodGet, "/api/v1/stats", nil) @@ -683,6 +679,7 @@ func TestRenderStats(t *testing.T) { var stats map[string]interface{} err := json.Unmarshal(rr.Body.Bytes(), &stats) require.NoError(t, err) - assert.Empty(t, stats) + assert.Equal(t, float64(0), stats["CurrentlyConnected"]) + assert.Contains(t, stats, "Since") }) } |