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:
+3
-93
@@ -6,63 +6,21 @@ package mw
|
|||||||
import (
|
import (
|
||||||
"crussell/clock"
|
"crussell/clock"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
|
||||||
"net/http"
|
|
||||||
"sync"
|
|
||||||
"net"
|
"net"
|
||||||
|
"net/http"
|
||||||
"time"
|
"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 {
|
func NewRateLimiter(limit int, window time.Duration) *RateLimiter {
|
||||||
rl := &RateLimiter{
|
rl := &RateLimiter{
|
||||||
requests: make(map[string][]time.Time),
|
requests: make(map[string][]time.Time),
|
||||||
limit: limit,
|
limit: limit,
|
||||||
window: window,
|
window: window,
|
||||||
}
|
}
|
||||||
// Cleanup old entries periodically
|
registerLimiter(rl)
|
||||||
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()
|
|
||||||
}()
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
return 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 {
|
func (rl *RateLimiter) Allow(key string) bool {
|
||||||
rl.mu.Lock()
|
rl.mu.Lock()
|
||||||
defer rl.mu.Unlock()
|
defer rl.mu.Unlock()
|
||||||
@@ -85,55 +43,10 @@ func (rl *RateLimiter) Allow(key string) bool {
|
|||||||
return true
|
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 {
|
func NewProgressiveRateLimiter() *ProgressiveRateLimiter {
|
||||||
prl := &ProgressiveRateLimiter{
|
return &ProgressiveRateLimiter{
|
||||||
requests: make(map[string]*ipProgressiveState),
|
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.
|
// 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 {
|
func ProgressiveRateLimit(next http.Handler) http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
ip := r.Header.Get("CF-Connecting-IP")
|
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)
|
limiter := NewRateLimiter(limit, window)
|
||||||
return func(next http.Handler) http.Handler {
|
return func(next http.Handler) http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
// Get client IP
|
|
||||||
ip := r.Header.Get("CF-Connecting-IP")
|
ip := r.Header.Get("CF-Connecting-IP")
|
||||||
if ip == "" {
|
if ip == "" {
|
||||||
ip, _, _ = net.SplitHostPort(r.RemoteAddr)
|
ip, _, _ = net.SplitHostPort(r.RemoteAddr)
|
||||||
|
|||||||
+26
-27
@@ -5,10 +5,35 @@ package mw
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"sync"
|
|
||||||
"time"
|
"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 {
|
func RateLimit(limit int, window time.Duration) func(http.Handler) http.Handler {
|
||||||
return func(next http.Handler) http.Handler {
|
return func(next http.Handler) http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
@@ -4,115 +4,13 @@
|
|||||||
package mw
|
package mw
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"sync"
|
"context"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"crussell/clock"
|
"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
|
// TestProgressiveRateLimiter_CleanupRemovesStaleEntries verifies that
|
||||||
// IPs with no activity for 60s are cleaned up.
|
// IPs with no activity for 60s are cleaned up.
|
||||||
func TestProgressiveRateLimiter_CleanupRemovesStaleEntries(t *testing.T) {
|
func TestProgressiveRateLimiter_CleanupRemovesStaleEntries(t *testing.T) {
|
||||||
@@ -124,7 +22,7 @@ func TestProgressiveRateLimiter_CleanupRemovesStaleEntries(t *testing.T) {
|
|||||||
}
|
}
|
||||||
prl.mu.Unlock()
|
prl.mu.Unlock()
|
||||||
|
|
||||||
prl.cleanup()
|
prl.Cleanup()
|
||||||
|
|
||||||
prl.mu.RLock()
|
prl.mu.RLock()
|
||||||
_, exists := prl.requests["stale-ip"]
|
_, 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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user