package ratelimit_test import ( "net/netip" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/ratelimit" ) // limit is the limit the tests set. const limit = 3 // The windows, as Count names them. const ( minute = "minute" hour = "hour" ) func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) { t.Parallel() for _, tc := range []struct { window string limits ratelimit.Limits length time.Duration }{ {minute, ratelimit.Limits{PerMinute: limit}, time.Minute}, {hour, ratelimit.Limits{PerHour: limit}, time.Hour}, {"day", ratelimit.Limits{PerDay: limit}, 24 * time.Hour}, } { t.Run(tc.window, func(t *testing.T) { t.Parallel() limiter := ratelimit.New(tc.limits) client := netip.MustParsePrefix("203.0.113.9/32") start := midnight() quarter := tc.length / 4 for range limit { wantCount(t, limiter, client, start, "") } wantCount(t, limiter, client, start, tc.window) // A quarter into the next bucket, the window still covers three // quarters of the bucket before, with its four requests: 3 + 1 // is over the limit. wantCount(t, limiter, client, start.Add(tc.length+quarter), tc.window) // Three quarters into it, a quarter: 1 + 2 is within. wantCount(t, limiter, client, start.Add(tc.length+3*quarter), "") }) } } func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) { t.Parallel() limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit}) client := netip.MustParsePrefix("203.0.113.9/32") start := midnight() for range limit { _, over := limiter.Count(client, start) if over { t.Fatal("a request within the limit is over it") } } // Over both limits; the minute's is named, with the four requests. hit, over := limiter.Count(client, start) want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1} if !over || hit != want { t.Errorf("request over the limit gives %+v and %t, want %+v and true", hit, over, want) } } func TestResetSetsTheCountsBackToZero(t *testing.T) { t.Parallel() limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit}) client := netip.MustParsePrefix("203.0.113.9/32") start := midnight() for range limit { wantCount(t, limiter, client, start, "") } wantCount(t, limiter, client, start, minute) limiter.Reset(client) // At the same moment, the client has its whole allowance again. for range limit { wantCount(t, limiter, client, start, "") } wantCount(t, limiter, client, start, minute) } func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) { t.Parallel() limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}) client := netip.MustParsePrefix("203.0.113.9/32") start := midnight() for range limit { wantCount(t, limiter, client, start, "") } wantCount(t, limiter, client, start, hour) // No request in the whole next bucket, so a quarter into the one after // it the window covers none of the four requests: 1 is within the // limit. Were they counted as the bucket before, 3 + 1 would be over. wantCount(t, limiter, client, start.Add(2*time.Hour+time.Hour/4), "") } func TestRefusedRequestsCount(t *testing.T) { t.Parallel() limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: 2 * limit}) refused := netip.MustParsePrefix("203.0.113.9/32") within := netip.MustParsePrefix("203.0.113.10/32") start := midnight() for range limit { wantCount(t, limiter, refused, start, "") wantCount(t, limiter, within, start, "") } for range limit { wantCount(t, limiter, refused, start, minute) } // Half a minute into the next bucket the window covers half of the // bucket before: 3 + 1 is over the minute's limit for the client // whose three refused requests count, and 1.5 + 1 within it for the // other. The first is over the hour's limit too, and the shorter // window is named. halfway := start.Add(time.Minute + time.Minute/2) wantCount(t, limiter, refused, halfway, minute) wantCount(t, limiter, within, halfway, "") // The refused requests count in the hour as well: 6 + 1 + 1 is over // its limit, and 3 + 1 + 1 within it. later := start.Add(10 * time.Minute) wantCount(t, limiter, refused, later, hour) wantCount(t, limiter, within, later, "") } func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) { t.Parallel() limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}) client := netip.MustParsePrefix("203.0.113.9/32") start := midnight() for range limit { wantCount(t, limiter, client, start, "") } // A concurrent request dated a moment before the bucket under way, but // counted after it began, is counted in it: 3 + 1 is over the limit. wantCount(t, limiter, client, start.Add(-time.Millisecond), minute) } func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) { t.Parallel() limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}) client := netip.MustParsePrefix("203.0.113.9/32") start := midnight() for range limit { wantCount(t, limiter, client, start, "") } // Half an hour into the next bucket: 3 / 2 + 1 is within the limit. wantCount(t, limiter, client, start.Add(time.Hour+time.Hour/2), "") // The clock is set back an hour. Counted in the bucket under way, the // next request would find the bucket before it at full weight, 3 + 2, // over the limit until the clock caught up. The buckets start afresh // instead, and the client is refused only past the limit again. setBack := start.Add(time.Hour / 2) for range limit { wantCount(t, limiter, client, setBack, "") } wantCount(t, limiter, client, setBack, hour) } func TestKeepsAtMost20000ClientsDroppingTheLeastRecentlySeen(t *testing.T) { t.Parallel() const maxClients = 20000 limiter := ratelimit.New(ratelimit.Limits{PerMinute: 1}) now := midnight() clients := make([]netip.Prefix, maxClients+1) addr := netip.MustParseAddr("10.0.0.0") for i := range clients { clients[i] = netip.PrefixFrom(addr, addr.BitLen()) addr = addr.Next() } for _, client := range clients[:maxClients] { wantCount(t, limiter, client, now, "") } // The first client is seen again: its second request is over the // limit of one, so it is still counted. wantCount(t, limiter, clients[0], now, minute) // One client more drops the least recently seen, the second, which // starts afresh, while the first is kept. wantCount(t, limiter, clients[maxClients], now, "") wantCount(t, limiter, clients[1], now, "") wantCount(t, limiter, clients[0], now, minute) } // midnight is the start of a bucket in every window. func midnight() time.Time { return time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC) } // wantCount counts a request from client at now, and checks the window // whose limit it goes over, "" for none. func wantCount( t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix, now time.Time, want string, ) { t.Helper() hit, _ := limiter.Count(client, now) if hit.Window != want { t.Errorf("request from %s at %s is over %q, want %q", client, now.Format(time.RFC3339), hit.Window, want) } }