diff --git a/internal/server/limiter.go b/internal/server/limiter.go index b94cdd8..9250d5a 100644 --- a/internal/server/limiter.go +++ b/internal/server/limiter.go @@ -41,7 +41,14 @@ func (l *Limiter) Allow(key string) (remaining int64, ok bool, retryAfter time.D shard.mu.Lock() defer shard.mu.Unlock() if len(shard.buckets) >= l.maxBuckets && shard.buckets[key] == nil { - l.evictShard(shard, now) + l.trimShard(shard, now.Add(-l.window)) + if len(shard.buckets) >= l.maxBuckets { + // Fail closed: evicting a live bucket would let an attacker + // mint fresh keys until their own (or a victim's) full bucket + // is dropped, resetting its counter and defeating the limit. + // Deny unknown keys while the shard is full instead. + return 0, false, l.window + } } hits := shard.buckets[key] cutoff := now.Add(-l.window) @@ -73,23 +80,7 @@ func (l *Limiter) sweep() { } } -// evictShard makes room in a full shard: expired buckets first, then -// arbitrary victims if the attack is still filling the window. -func (l *Limiter) evictShard(shard *limShard, now time.Time) { - l.trimShard(shard, now.Add(-l.window)) - for len(shard.buckets) >= l.maxBuckets { - victimized := false - for k := range shard.buckets { - delete(shard.buckets, k) - victimized = true - break - } - if !victimized { - break - } - } -} - +// trimShard drops buckets whose hits have all expired past cutoff. func (l *Limiter) trimShard(shard *limShard, cutoff time.Time) { for k, hits := range shard.buckets { j := 0 diff --git a/internal/server/limiter_test.go b/internal/server/limiter_test.go index bc18005..79a9f78 100644 --- a/internal/server/limiter_test.go +++ b/internal/server/limiter_test.go @@ -4,6 +4,7 @@ import ( "fmt" "net/http/httptest" "testing" + "time" ) func TestClientKeyIgnoresBogusXFF(t *testing.T) { @@ -37,3 +38,37 @@ func TestLimiterBucketsBounded(t *testing.T) { t.Errorf("buckets = %d, want <= %d", total, maxBucketsPerShard*64) } } + +func TestLimiterFullShardFailsClosed(t *testing.T) { + l := &Limiter{limit: 10, window: time.Hour, maxBuckets: 1} + for i := range l.shards { + l.shards[i].buckets = make(map[string][]time.Time) + } + first := "client-a" + peer := "" + for i := 0; i < 10000 && peer == ""; i++ { + if k := fmt.Sprintf("k%d", i); shardIndex(k) == shardIndex(first) && k != first { + peer = k + } + } + if peer == "" { + t.Fatal("no key found hashing to the same shard") + } + + if _, ok, _ := l.Allow(first); !ok { + t.Fatal("first key should be admitted") + } + // Shard is full (maxBuckets=1). The unknown peer key must be denied + // rather than evicting the live bucket, which would reset first's + // counter. + _, ok, retry := l.Allow(peer) + if ok { + t.Fatal("unknown key should be denied when shard is full") + } + if retry <= 0 { + t.Fatalf("retryAfter = %v; want > 0", retry) + } + if _, ok, _ := l.Allow(first); !ok { + t.Fatal("existing key's bucket must survive peer admission attempts") + } +}