From 6064a06b7dd1aaea631a141c68bc0129b757f768 Mon Sep 17 00:00:00 2001 From: Stephen Adamson Date: Tue, 7 Jul 2026 00:09:43 +0100 Subject: [PATCH] refactor: extract shared rate limiter types into ratelimit_shared.go Move RateLimiter, ProgressiveRateLimiter, and cleanup functions to ratelimit_shared.go. Add dev and prod rate limiter tests. Remove inline cleanup goroutines. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus --- backend/mw/ratelimit.go | 96 +----------- backend/mw/ratelimit_dev.go | 53 ++++--- backend/mw/ratelimit_dev_test.go | 193 ++++++++++++++++++++++++ backend/mw/ratelimit_prod_test.go | 113 ++++++++++++++ backend/mw/ratelimit_shared.go | 104 +++++++++++++ backend/mw/ratelimit_test.go | 240 +++++++++++++++++------------- 6 files changed, 574 insertions(+), 225 deletions(-) create mode 100644 backend/mw/ratelimit_dev_test.go create mode 100644 backend/mw/ratelimit_prod_test.go create mode 100644 backend/mw/ratelimit_shared.go diff --git a/backend/mw/ratelimit.go b/backend/mw/ratelimit.go index dc0c357..c1a7efe 100644 --- a/backend/mw/ratelimit.go +++ b/backend/mw/ratelimit.go @@ -6,63 +6,21 @@ package mw import ( "crussell/clock" "fmt" - "log" - "net/http" - "sync" "net" + "net/http" "time" ) -// RateLimiter implements a simple in-memory rate limiter -type RateLimiter struct { - requests map[string][]time.Time - mu sync.RWMutex - limit int - window time.Duration -} - func NewRateLimiter(limit int, window time.Duration) *RateLimiter { rl := &RateLimiter{ requests: make(map[string][]time.Time), limit: limit, window: window, } - // Cleanup old entries periodically - go func() { - for { - func() { - defer func() { - if r := recover(); r != nil { - log.Printf("Panic recovered in rate limiter cleanup: %v", r) - } - }() - time.Sleep(window) - rl.cleanup() - }() - } - }() + registerLimiter(rl) return rl } -func (rl *RateLimiter) cleanup() { - rl.mu.Lock() - defer rl.mu.Unlock() - now := clock.Now() - for key, times := range rl.requests { - var valid []time.Time - for _, t := range times { - if now.Sub(t) < rl.window { - valid = append(valid, t) - } - } - if len(valid) == 0 { - delete(rl.requests, key) - } else { - rl.requests[key] = valid - } - } -} - func (rl *RateLimiter) Allow(key string) bool { rl.mu.Lock() defer rl.mu.Unlock() @@ -85,55 +43,10 @@ func (rl *RateLimiter) Allow(key string) bool { return true } -// ProgressiveRateLimiter implements per-IP rate limiting with increasing backoff -// Designed for bot-spam prevention across accounts (not account-specific lockout) -type ProgressiveRateLimiter struct { - requests map[string]*ipProgressiveState - mu sync.RWMutex -} - -type ipProgressiveState struct { - // Timestamps of all requests within the tracking window - timestamps []time.Time -} - func NewProgressiveRateLimiter() *ProgressiveRateLimiter { - prl := &ProgressiveRateLimiter{ + return &ProgressiveRateLimiter{ requests: make(map[string]*ipProgressiveState), } - go func() { - for { - func() { - defer func() { - if r := recover(); r != nil { - log.Printf("Panic recovered in progressive rate limiter cleanup: %v", r) - } - }() - time.Sleep(30 * time.Second) - prl.cleanup() - }() - } - }() - return prl -} - -func (prl *ProgressiveRateLimiter) cleanup() { - prl.mu.Lock() - defer prl.mu.Unlock() - cutoff := clock.Now().Add(-60 * time.Second) - for ip, state := range prl.requests { - var valid []time.Time - for _, t := range state.timestamps { - if t.After(cutoff) { - valid = append(valid, t) - } - } - if len(valid) == 0 { - delete(prl.requests, ip) - } else { - state.timestamps = valid - } - } } // Check returns the delay in milliseconds. Returns 0 if no delay needed. @@ -190,8 +103,6 @@ func (prl *ProgressiveRateLimiter) Check(ip string) (delayMs int) { } } -var globalProgressiveLimiter = NewProgressiveRateLimiter() - func ProgressiveRateLimit(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ip := r.Header.Get("CF-Connecting-IP") @@ -217,7 +128,6 @@ func RateLimit(limit int, window time.Duration) func(http.Handler) http.Handler limiter := NewRateLimiter(limit, window) return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Get client IP ip := r.Header.Get("CF-Connecting-IP") if ip == "" { ip, _, _ = net.SplitHostPort(r.RemoteAddr) diff --git a/backend/mw/ratelimit_dev.go b/backend/mw/ratelimit_dev.go index 3efc47f..634704a 100644 --- a/backend/mw/ratelimit_dev.go +++ b/backend/mw/ratelimit_dev.go @@ -5,10 +5,35 @@ package mw import ( "net/http" - "sync" "time" ) +func NewRateLimiter(limit int, window time.Duration) *RateLimiter { + rl := &RateLimiter{ + requests: make(map[string][]time.Time), + limit: limit, + window: window, + } + registerLimiter(rl) + return rl +} + +func (rl *RateLimiter) Allow(key string) bool { return true } + +func NewProgressiveRateLimiter() *ProgressiveRateLimiter { + return &ProgressiveRateLimiter{ + requests: make(map[string]*ipProgressiveState), + } +} + +func (prl *ProgressiveRateLimiter) Check(ip string) int { return 0 } + +func ProgressiveRateLimit(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + next.ServeHTTP(w, r) + }) +} + func RateLimit(limit int, window time.Duration) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -16,29 +41,3 @@ func RateLimit(limit int, window time.Duration) func(http.Handler) http.Handler }) } } - -type ProgressiveRateLimiter struct { - requests map[string]*ipProgressiveState - mu sync.RWMutex -} - -type ipProgressiveState struct { - timestamps []time.Time -} - -func NewProgressiveRateLimiter() *ProgressiveRateLimiter { - return &ProgressiveRateLimiter{ - requests: make(map[string]*ipProgressiveState), - } -} - -func (prl *ProgressiveRateLimiter) cleanup() {} -func (prl *ProgressiveRateLimiter) Check(ip string) int { return 0 } - -var globalProgressiveLimiter = NewProgressiveRateLimiter() - -func ProgressiveRateLimit(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - next.ServeHTTP(w, r) - }) -} diff --git a/backend/mw/ratelimit_dev_test.go b/backend/mw/ratelimit_dev_test.go new file mode 100644 index 0000000..1d47c0e --- /dev/null +++ b/backend/mw/ratelimit_dev_test.go @@ -0,0 +1,193 @@ +//go:build test && dev +// +build test,dev + +package mw + +import ( + "context" + "testing" + "time" + + "crussell/clock" +) + +// ============================================================ +// Dev Stub Behavioral Tests +// ============================================================ + +// TestDevRateLimiter_AllowAlwaysTrue verifies the dev stub never throttles. +func TestDevRateLimiter_AllowAlwaysTrue(t *testing.T) { + rl := NewRateLimiter(10, time.Minute) + + if !rl.Allow("any-key") { + t.Error("expected Allow to return true for dev stub") + } + if !rl.Allow("another-key") { + t.Error("expected Allow to return true for dev stub (different key)") + } + // Even an empty key should pass + if !rl.Allow("") { + t.Error("expected Allow to return true for empty key in dev stub") + } +} + +// TestDevProgressiveRateLimiter_CheckAlwaysZero verifies the dev stub never delays. +func TestDevProgressiveRateLimiter_CheckAlwaysZero(t *testing.T) { + prl := NewProgressiveRateLimiter() + + for i := 0; i < 200; i++ { + delay := prl.Check("10.0.0.1") + if delay != 0 { + t.Errorf("expected 0 delay for dev stub (request %d), got %d", i+1, delay) + } + } + + // Different IP also returns 0 + if delay := prl.Check("10.0.0.2"); delay != 0 { + t.Errorf("expected 0 delay for different IP in dev stub, got %d", delay) + } +} + +func TestProgressiveRateLimiter_CleanupRemovesStaleEntries(t *testing.T) { + prl := NewProgressiveRateLimiter() + + prl.mu.Lock() + prl.requests["stale-ip"] = &ipProgressiveState{ + timestamps: []time.Time{clock.Now().Add(-120 * time.Second)}, + } + prl.mu.Unlock() + + prl.Cleanup() + + prl.mu.RLock() + _, exists := prl.requests["stale-ip"] + prl.mu.RUnlock() + if exists { + t.Error("expected stale IP to be cleaned up") + } +} + +func TestRateLimiter_Cleanup(t *testing.T) { + rl := NewRateLimiter(10, time.Minute) + + rl.mu.Lock() + rl.requests["stale-key"] = []time.Time{clock.Now().Add(-5 * time.Minute)} + rl.requests["fresh-key"] = []time.Time{clock.Now()} + rl.mu.Unlock() + + rl.Cleanup() + + rl.mu.RLock() + _, staleExists := rl.requests["stale-key"] + _, freshExists := rl.requests["fresh-key"] + rl.mu.RUnlock() + + if staleExists { + t.Error("expected stale-key to be removed") + } + if !freshExists { + t.Error("expected fresh-key to be preserved") + } +} + +func TestRateLimiter_Cleanup_EmptyMap(t *testing.T) { + rl := NewRateLimiter(10, time.Minute) + + rl.mu.Lock() + rl.requests = make(map[string][]time.Time) + rl.mu.Unlock() + + rl.Cleanup() + + rl.mu.RLock() + count := len(rl.requests) + rl.mu.RUnlock() + if count != 0 { + t.Errorf("expected empty map, got %d entries", count) + } +} + +func TestCleanupAllRateLimiters(t *testing.T) { + registeredLimitersMu.Lock() + saved := registeredLimiters + registeredLimitersMu.Unlock() + defer func() { + registeredLimitersMu.Lock() + registeredLimiters = saved + registeredLimitersMu.Unlock() + }() + + rl1 := NewRateLimiter(10, time.Minute) + rl2 := NewRateLimiter(20, time.Minute) + + rl1.mu.Lock() + rl1.requests["rl1-stale"] = []time.Time{clock.Now().Add(-5 * time.Minute)} + rl1.mu.Unlock() + + rl2.mu.Lock() + rl2.requests["rl2-stale"] = []time.Time{clock.Now().Add(-5 * time.Minute)} + rl2.mu.Unlock() + + _, err := CleanupAllRateLimiters(context.Background()) + if err != nil { + t.Errorf("expected nil error, got %v", err) + } + + rl1.mu.RLock() + _, rl1Stale := rl1.requests["rl1-stale"] + rl1.mu.RUnlock() + + rl2.mu.RLock() + _, rl2Stale := rl2.requests["rl2-stale"] + rl2.mu.RUnlock() + + if rl1Stale { + t.Error("expected rl1-stale to be removed") + } + if rl2Stale { + t.Error("expected rl2-stale to be removed") + } +} + +func TestCleanupAllRateLimiters_Empty(t *testing.T) { + registeredLimitersMu.Lock() + saved := registeredLimiters + registeredLimiters = nil + registeredLimitersMu.Unlock() + defer func() { + registeredLimitersMu.Lock() + registeredLimiters = saved + registeredLimitersMu.Unlock() + }() + + _, err := CleanupAllRateLimiters(context.Background()) + if err != nil { + t.Errorf("expected nil error, got %v", err) + } +} + +func TestCleanupProgressiveRateLimiter(t *testing.T) { + globalProgressiveLimiter.mu.Lock() + saved := globalProgressiveLimiter.requests + globalProgressiveLimiter.requests = map[string]*ipProgressiveState{ + "global-stale": {timestamps: []time.Time{clock.Now().Add(-120 * time.Second)}}, + } + globalProgressiveLimiter.mu.Unlock() + defer func() { + globalProgressiveLimiter.mu.Lock() + globalProgressiveLimiter.requests = saved + globalProgressiveLimiter.mu.Unlock() + }() + + _, err := CleanupProgressiveRateLimiter(context.Background()) + if err != nil { + t.Errorf("expected nil error, got %v", err) + } + + globalProgressiveLimiter.mu.RLock() + _, exists := globalProgressiveLimiter.requests["global-stale"] + globalProgressiveLimiter.mu.RUnlock() + if exists { + t.Error("expected global-stale to be removed") + } +} diff --git a/backend/mw/ratelimit_prod_test.go b/backend/mw/ratelimit_prod_test.go new file mode 100644 index 0000000..c4ed514 --- /dev/null +++ b/backend/mw/ratelimit_prod_test.go @@ -0,0 +1,113 @@ +//go:build test && !dev +// +build test,!dev + +package mw + +import ( + "testing" + "time" + + "crussell/clock" +) + +// TestProgressiveRateLimiter_SingleRequest verifies no delay for first request. +func TestProgressiveRateLimiter_SingleRequest(t *testing.T) { + prl := NewProgressiveRateLimiter() + delay := prl.Check("192.168.1.1") + if delay != 0 { + t.Errorf("expected 0 delay for first request, got %d", delay) + } +} + +// TestProgressiveRateLimiter_BurstAllowsPageLoad verifies 15 requests in 5s are OK. +func TestProgressiveRateLimiter_BurstAllowsPageLoad(t *testing.T) { + prl := NewProgressiveRateLimiter() + ip := "192.168.1.1" + + for i := 0; i < 15; i++ { + delay := prl.Check(ip) + if delay != 0 { + t.Errorf("expected 0 delay for request %d (within burst), got %d", i+1, delay) + } + } +} + +// TestProgressiveRateLimiter_ExcessBurstDelays verifies 35+ requests in 5s get delayed. +func TestProgressiveRateLimiter_ExcessBurstDelays(t *testing.T) { + prl := NewProgressiveRateLimiter() + ip := "192.168.1.1" + + delayed := false + for i := 0; i < 35; i++ { + delay := prl.Check(ip) + if delay > 0 { + delayed = true + } + } + + if !delayed { + t.Error("expected at least one delay after 35 burst requests") + } +} + +// TestProgressiveRateLimiter_SustainedAllowsNormal verifies 30 requests +// spread over 60 seconds are not delayed. +func TestProgressiveRateLimiter_SustainedAllowsNormal(t *testing.T) { + prl := NewProgressiveRateLimiter() + ip := "192.168.1.2" + + for i := 0; i < 30; i++ { + prl.mu.Lock() + state, exists := prl.requests[ip] + if !exists { + prl.requests[ip] = &ipProgressiveState{ + timestamps: []time.Time{clock.Now().Add(-time.Duration(60-i*2) * time.Second)}, + } + prl.mu.Unlock() + continue + } + state.timestamps = append(state.timestamps, clock.Now().Add(-time.Duration(60-i*2)*time.Second)) + prl.mu.Unlock() + } + + delay := prl.Check(ip) + if delay != 0 { + t.Errorf("expected 0 delay for 30 sustained requests over 60s, got %d", delay) + } +} + +// TestProgressiveRateLimiter_ExcessSustainedDelays verifies 130+ requests +// in 60s triggers delay. +func TestProgressiveRateLimiter_ExcessSustainedDelays(t *testing.T) { + prl := NewProgressiveRateLimiter() + ip := "192.168.1.3" + + now := clock.Now() + prl.mu.Lock() + state := &ipProgressiveState{timestamps: make([]time.Time, 130)} + for i := 0; i < 130; i++ { + state.timestamps[i] = now.Add(-time.Duration(60-i/3) * time.Second) + } + prl.requests[ip] = state + prl.mu.Unlock() + + delay := prl.Check(ip) + if delay == 0 { + t.Error("expected delay > 0 for 130 sustained requests") + } +} + +// TestProgressiveRateLimiter_DifferentIPs verifies rate limiter +// tracks IPs independently. +func TestProgressiveRateLimiter_DifferentIPs(t *testing.T) { + prl := NewProgressiveRateLimiter() + + for i := 0; i < 40; i++ { + prl.Check("10.0.0.1") + } + + delay := prl.Check("10.0.0.2") + if delay != 0 { + t.Errorf("expected 0 delay for separate IP, got %d", delay) + } +} diff --git a/backend/mw/ratelimit_shared.go b/backend/mw/ratelimit_shared.go new file mode 100644 index 0000000..37f4b27 --- /dev/null +++ b/backend/mw/ratelimit_shared.go @@ -0,0 +1,104 @@ +// Shared types, registration, and cleanup for RateLimiter + ProgressiveRateLimiter. +// Used by both production (!dev) and dev (dev) builds — keep tag-free. + +package mw + +import ( + "context" + "crussell/clock" + "sync" + "time" +) + +// RateLimiter implements a simple in-memory rate limiter +type RateLimiter struct { + requests map[string][]time.Time + mu sync.RWMutex + limit int + window time.Duration +} + +// ProgressiveRateLimiter implements per-IP rate limiting with increasing backoff. +// Designed for bot-spam prevention across accounts (not account-specific lockout). +type ProgressiveRateLimiter struct { + requests map[string]*ipProgressiveState + mu sync.RWMutex +} + +type ipProgressiveState struct { + // Timestamps of all requests within the tracking window + timestamps []time.Time +} + +var ( + registeredLimiters []*RateLimiter + registeredLimitersMu sync.Mutex +) + +func registerLimiter(rl *RateLimiter) { + registeredLimitersMu.Lock() + registeredLimiters = append(registeredLimiters, rl) + registeredLimitersMu.Unlock() +} + +// CleanupAllRateLimiters runs Cleanup on every registered RateLimiter. +// Called by the centralised jobs scheduler. +func CleanupAllRateLimiters(ctx context.Context) (int, error) { + registeredLimitersMu.Lock() + limiters := make([]*RateLimiter, len(registeredLimiters)) + copy(limiters, registeredLimiters) + registeredLimitersMu.Unlock() + for _, l := range limiters { + l.Cleanup() + } + return 0, nil +} + +// CleanupProgressiveRateLimiter runs Cleanup on the global progressive rate limiter. +// Called by the centralised jobs scheduler. +func CleanupProgressiveRateLimiter(ctx context.Context) (int, error) { + globalProgressiveLimiter.Cleanup() + return 0, nil +} + +// Cleanup removes expired entries from the rate limiter map. +func (rl *RateLimiter) Cleanup() { + rl.mu.Lock() + defer rl.mu.Unlock() + now := clock.Now() + for key, times := range rl.requests { + var valid []time.Time + for _, t := range times { + if now.Sub(t) < rl.window { + valid = append(valid, t) + } + } + if len(valid) == 0 { + delete(rl.requests, key) + } else { + rl.requests[key] = valid + } + } +} + +// Cleanup removes expired entries from the progressive rate limiter map. +func (prl *ProgressiveRateLimiter) Cleanup() { + prl.mu.Lock() + defer prl.mu.Unlock() + cutoff := clock.Now().Add(-60 * time.Second) + for ip, state := range prl.requests { + var valid []time.Time + for _, t := range state.timestamps { + if t.After(cutoff) { + valid = append(valid, t) + } + } + if len(valid) == 0 { + delete(prl.requests, ip) + } else { + state.timestamps = valid + } + } +} + +var globalProgressiveLimiter = NewProgressiveRateLimiter() diff --git a/backend/mw/ratelimit_test.go b/backend/mw/ratelimit_test.go index 99fe07d..3db7543 100644 --- a/backend/mw/ratelimit_test.go +++ b/backend/mw/ratelimit_test.go @@ -4,115 +4,13 @@ package mw import ( - "sync" + "context" "testing" "time" "crussell/clock" ) -// TestProgressiveRateLimiter_SingleRequest verifies no delay for first request. -func TestProgressiveRateLimiter_SingleRequest(t *testing.T) { - prl := NewProgressiveRateLimiter() - delay := prl.Check("192.168.1.1") - if delay != 0 { - t.Errorf("expected 0 delay for first request, got %d", delay) - } -} - -// TestProgressiveRateLimiter_BurstAllowsPageLoad verifies 15 requests in 5s are OK. -func TestProgressiveRateLimiter_BurstAllowsPageLoad(t *testing.T) { - prl := NewProgressiveRateLimiter() - ip := "192.168.1.1" - - for i := 0; i < 15; i++ { - delay := prl.Check(ip) - if delay != 0 { - t.Errorf("expected 0 delay for request %d (within burst), got %d", i+1, delay) - } - } -} - -// TestProgressiveRateLimiter_ExcessBurstDelays verifies 35+ requests in 5s get delayed. -func TestProgressiveRateLimiter_ExcessBurstDelays(t *testing.T) { - prl := NewProgressiveRateLimiter() - ip := "192.168.1.1" - - delayed := false - for i := 0; i < 35; i++ { - delay := prl.Check(ip) - if delay > 0 { - delayed = true - } - } - - if !delayed { - t.Error("expected at least one delay after 35 burst requests") - } -} - -// TestProgressiveRateLimiter_SustainedAllowsNormal verifies 30 requests -// spread over 60 seconds are not delayed. -func TestProgressiveRateLimiter_SustainedAllowsNormal(t *testing.T) { - prl := NewProgressiveRateLimiter() - ip := "192.168.1.2" - - for i := 0; i < 30; i++ { - prl.mu.Lock() - state, exists := prl.requests[ip] - if !exists { - prl.requests[ip] = &ipProgressiveState{ - timestamps: []time.Time{clock.Now().Add(-time.Duration(60-i*2) * time.Second)}, - } - prl.mu.Unlock() - continue - } - state.timestamps = append(state.timestamps, clock.Now().Add(-time.Duration(60-i*2)*time.Second)) - prl.mu.Unlock() - } - - delay := prl.Check(ip) - if delay != 0 { - t.Errorf("expected 0 delay for 30 sustained requests over 60s, got %d", delay) - } -} - -// TestProgressiveRateLimiter_ExcessSustainedDelays verifies 130+ requests -// in 60s triggers delay. -func TestProgressiveRateLimiter_ExcessSustainedDelays(t *testing.T) { - prl := NewProgressiveRateLimiter() - ip := "192.168.1.3" - - now := clock.Now() - prl.mu.Lock() - state := &ipProgressiveState{timestamps: make([]time.Time, 130)} - for i := 0; i < 130; i++ { - state.timestamps[i] = now.Add(-time.Duration(60-i/3) * time.Second) - } - prl.requests[ip] = state - prl.mu.Unlock() - - delay := prl.Check(ip) - if delay == 0 { - t.Error("expected delay > 0 for 130 sustained requests") - } -} - -// TestProgressiveRateLimiter_DifferentIPs verifies rate limiter -// tracks IPs independently. -func TestProgressiveRateLimiter_DifferentIPs(t *testing.T) { - prl := NewProgressiveRateLimiter() - - for i := 0; i < 40; i++ { - prl.Check("10.0.0.1") - } - - delay := prl.Check("10.0.0.2") - if delay != 0 { - t.Errorf("expected 0 delay for separate IP, got %d", delay) - } -} - // TestProgressiveRateLimiter_CleanupRemovesStaleEntries verifies that // IPs with no activity for 60s are cleaned up. func TestProgressiveRateLimiter_CleanupRemovesStaleEntries(t *testing.T) { @@ -124,7 +22,7 @@ func TestProgressiveRateLimiter_CleanupRemovesStaleEntries(t *testing.T) { } prl.mu.Unlock() - prl.cleanup() + prl.Cleanup() prl.mu.RLock() _, exists := prl.requests["stale-ip"] @@ -134,4 +32,136 @@ func TestProgressiveRateLimiter_CleanupRemovesStaleEntries(t *testing.T) { } } -var _ = sync.Mutex{} +// ============================================================ +// RateLimiter Cleanup Tests +// ============================================================ + +// TestRateLimiter_Cleanup verifies that stale entries are removed. +func TestRateLimiter_Cleanup(t *testing.T) { + rl := NewRateLimiter(10, time.Minute) + + rl.mu.Lock() + rl.requests["stale-key"] = []time.Time{clock.Now().Add(-5 * time.Minute)} + rl.requests["fresh-key"] = []time.Time{clock.Now()} + rl.mu.Unlock() + + rl.Cleanup() + + rl.mu.RLock() + _, staleExists := rl.requests["stale-key"] + _, freshExists := rl.requests["fresh-key"] + rl.mu.RUnlock() + + if staleExists { + t.Error("expected stale-key to be removed") + } + if !freshExists { + t.Error("expected fresh-key to be preserved") + } +} + +// TestRateLimiter_Cleanup_EmptyMap handles an empty requests map. +func TestRateLimiter_Cleanup_EmptyMap(t *testing.T) { + rl := NewRateLimiter(10, time.Minute) + + rl.mu.Lock() + rl.requests = make(map[string][]time.Time) + rl.mu.Unlock() + + rl.Cleanup() + + rl.mu.RLock() + count := len(rl.requests) + rl.mu.RUnlock() + if count != 0 { + t.Errorf("expected empty map, got %d entries", count) + } +} + +// TestCleanupAllRateLimiters iterates all registered limiters. +func TestCleanupAllRateLimiters(t *testing.T) { + registeredLimitersMu.Lock() + saved := registeredLimiters + registeredLimitersMu.Unlock() + defer func() { + registeredLimitersMu.Lock() + registeredLimiters = saved + registeredLimitersMu.Unlock() + }() + + rl1 := NewRateLimiter(10, time.Minute) + rl2 := NewRateLimiter(20, time.Minute) + + rl1.mu.Lock() + rl1.requests["rl1-stale"] = []time.Time{clock.Now().Add(-5 * time.Minute)} + rl1.mu.Unlock() + + rl2.mu.Lock() + rl2.requests["rl2-stale"] = []time.Time{clock.Now().Add(-5 * time.Minute)} + rl2.mu.Unlock() + + _, err := CleanupAllRateLimiters(context.Background()) + if err != nil { + t.Errorf("expected nil error, got %v", err) + } + + rl1.mu.RLock() + _, rl1Stale := rl1.requests["rl1-stale"] + rl1.mu.RUnlock() + + rl2.mu.RLock() + _, rl2Stale := rl2.requests["rl2-stale"] + rl2.mu.RUnlock() + + if rl1Stale { + t.Error("expected rl1-stale to be removed") + } + if rl2Stale { + t.Error("expected rl2-stale to be removed") + } +} + +// TestCleanupAllRateLimiters_Empty does not panic with no registered limiters. +func TestCleanupAllRateLimiters_Empty(t *testing.T) { + registeredLimitersMu.Lock() + saved := registeredLimiters + registeredLimiters = nil + registeredLimitersMu.Unlock() + defer func() { + registeredLimitersMu.Lock() + registeredLimiters = saved + registeredLimitersMu.Unlock() + }() + + _, err := CleanupAllRateLimiters(context.Background()) + if err != nil { + t.Errorf("expected nil error, got %v", err) + } +} + +// TestCleanupProgressiveRateLimiter cleans the global limiter. +func TestCleanupProgressiveRateLimiter(t *testing.T) { + globalProgressiveLimiter.mu.Lock() + saved := globalProgressiveLimiter.requests + globalProgressiveLimiter.requests = map[string]*ipProgressiveState{ + "global-stale": {timestamps: []time.Time{clock.Now().Add(-120 * time.Second)}}, + } + globalProgressiveLimiter.mu.Unlock() + defer func() { + globalProgressiveLimiter.mu.Lock() + globalProgressiveLimiter.requests = saved + globalProgressiveLimiter.mu.Unlock() + }() + + _, err := CleanupProgressiveRateLimiter(context.Background()) + if err != nil { + t.Errorf("expected nil error, got %v", err) + } + + globalProgressiveLimiter.mu.RLock() + _, exists := globalProgressiveLimiter.requests["global-stale"] + globalProgressiveLimiter.mu.RUnlock() + if exists { + t.Error("expected global-stale to be removed") + } +}