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 <clio-agent@sisyphuslabs.ai>
This commit is contained in:
2026-07-07 00:09:43 +01:00
co-authored by Sisyphus
parent f78489ae00
commit 6064a06b7d
6 changed files with 574 additions and 225 deletions
+3 -93
View File
@@ -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)
+26 -27
View File
@@ -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)
})
}
+193
View File
@@ -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")
}
}
+113
View File
@@ -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)
}
}
+104
View File
@@ -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()
+135 -105
View File
@@ -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")
}
}