Files
krow_backend/go-api/internal/ratelimit/ratelimit_test.go
Aravind f2aa3b3ad8
Some checks failed
CI / fixture (push) Has been cancelled
CI / test (push) Has been cancelled
mcp connection
2026-09-22 10:58:02 +05:30

453 lines
14 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package ratelimit
import (
"context"
"strings"
"sync"
"testing"
"time"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
"github.com/krow/krow-backend/go-api/internal/testutil"
)
func newLimiter(t *testing.T) (*Limiter, *testutil.Harness, *time.Time) {
t.Helper()
h := testutil.New(t)
clock := time.Now().Truncate(time.Hour) // a clean window boundary
l := New(h.Pool).WithClock(func() time.Time { return clock })
return l, h, &clock
}
var testRule = Rule{Name: "test.rule", Limit: 3, Window: time.Minute}
func mustAllow(t *testing.T, l *Limiter, subject string) Decision {
t.Helper()
d, err := l.Allow(context.Background(), testRule, Subject(subject))
if err != nil {
t.Fatalf("Allow: %v", err)
}
return d
}
/* ── The basic contract ─────────────────────────────────────────────────── */
func TestUnderAtAndOverTheLimit(t *testing.T) {
l, _, _ := newLimiter(t)
// Under: each of the first three is allowed, and remaining counts down.
for i := 1; i <= testRule.Limit; i++ {
d := mustAllow(t, l, "alice")
if !d.Allowed {
t.Fatalf("request %d of %d was refused", i, testRule.Limit)
}
if want := testRule.Limit - i; d.Remaining != want {
t.Errorf("request %d: remaining = %d, want %d", i, d.Remaining, want)
}
}
// Over: the next one is refused and carries a usable Retry-After.
d := mustAllow(t, l, "alice")
if d.Allowed {
t.Fatal("the request past the limit was allowed")
}
if d.Remaining != 0 {
t.Errorf("remaining = %d, want 0", d.Remaining)
}
if d.RetryAfter <= 0 || d.RetryAfter > testRule.Window {
t.Errorf("RetryAfter = %v, want a positive interval no longer than the window", d.RetryAfter)
}
}
// A refused request is still counted. Otherwise a caller at their limit could
// keep probing for free, and the limit would be cheaper to test than to respect.
func TestRefusedRequestsStillCount(t *testing.T) {
l, h, _ := newLimiter(t)
for i := 0; i < testRule.Limit+5; i++ {
mustAllow(t, l, "bob")
}
var count int
if err := h.Pool.QueryRow(context.Background(),
`SELECT count FROM rate_limits WHERE bucket LIKE $1`, testRule.Name+":%").Scan(&count); err != nil {
t.Fatalf("read counter: %v", err)
}
if count != testRule.Limit+5 {
t.Errorf("count = %d, want %d — refusals must be counted too", count, testRule.Limit+5)
}
}
/* ── Windows ────────────────────────────────────────────────────────────── */
func TestTheWindowResets(t *testing.T) {
l, _, clock := newLimiter(t)
for i := 0; i < testRule.Limit; i++ {
mustAllow(t, l, "carol")
}
if mustAllow(t, l, "carol").Allowed {
t.Fatal("expected to be at the limit")
}
// Roll into the next window.
*clock = clock.Add(testRule.Window)
l.WithClock(func() time.Time { return *clock })
if d := mustAllow(t, l, "carol"); !d.Allowed {
t.Error("the limit did not reset at the window boundary")
} else if d.Remaining != testRule.Limit-1 {
t.Errorf("remaining = %d, want %d after a reset", d.Remaining, testRule.Limit-1)
}
}
// A new window is a new ROW, not a reset of an existing counter. That is what
// makes two instances rolling over simultaneously safe: neither clobbers the
// other's increments.
func TestANewWindowIsANewRow(t *testing.T) {
l, h, clock := newLimiter(t)
mustAllow(t, l, "dave")
*clock = clock.Add(testRule.Window)
l.WithClock(func() time.Time { return *clock })
mustAllow(t, l, "dave")
var rows int
if err := h.Pool.QueryRow(context.Background(),
`SELECT count(*) FROM rate_limits WHERE bucket LIKE $1`, testRule.Name+":%").Scan(&rows); err != nil {
t.Fatalf("count rows: %v", err)
}
if rows != 2 {
t.Errorf("%d rows, want 2 — each window must be its own row", rows)
}
}
/* ── Buckets are independent ────────────────────────────────────────────── */
func TestSubjectsAreIndependent(t *testing.T) {
l, _, _ := newLimiter(t)
// Exhaust one subject entirely.
for i := 0; i < testRule.Limit+2; i++ {
mustAllow(t, l, "user-a")
}
// A different subject must be untouched.
if d := mustAllow(t, l, "user-b"); !d.Allowed {
t.Error("one subject's limit affected another's")
}
if d := mustAllow(t, l, "org-a|user-a"); !d.Allowed {
t.Error("a compound subject collided with a simple one")
}
}
func TestRulesAreIndependent(t *testing.T) {
l, _, _ := newLimiter(t)
other := Rule{Name: "other.rule", Limit: 3, Window: time.Minute}
for i := 0; i < 5; i++ {
mustAllow(t, l, "shared")
}
d, err := l.Allow(context.Background(), other, Subject("shared"))
if err != nil {
t.Fatalf("Allow: %v", err)
}
if !d.Allowed {
t.Error("exhausting one rule exhausted another for the same subject")
}
}
/* ── No credential becomes a key ────────────────────────────────────────── */
// The property that matters most here: a bucket must never contain the thing it
// identifies. A token in this table is a token in every EXPLAIN, every slow
// query log and every backup.
func TestSubjectsAreHashedNotStored(t *testing.T) {
l, h, _ := newLimiter(t)
const secret = "a-very-secret-bearer-token-value"
mustAllow(t, l, secret)
var found int
if err := h.Pool.QueryRow(context.Background(),
`SELECT count(*) FROM rate_limits WHERE bucket LIKE '%' || $1 || '%'`,
secret).Scan(&found); err != nil {
t.Fatalf("scan: %v", err)
}
if found != 0 {
t.Error("the raw subject appears in the rate_limits table")
}
// And the hash is stable, or a caller would get a fresh budget per request.
if Subject(secret) != Subject(secret) {
t.Error("Subject is not deterministic")
}
if Subject(secret) == secret {
t.Error("Subject returned the raw value")
}
if len(Subject(secret)) != 32 {
t.Errorf("Subject length = %d, want 32", len(Subject(secret)))
}
}
/* ── Concurrency ────────────────────────────────────────────────────────── */
// The claim this package rests on: the count comes back from the statement that
// increments it, so concurrent callers cannot both read the same value and both
// proceed. Run with -race.
func TestConcurrentCallersDoNotLoseIncrements(t *testing.T) {
h := testutil.New(t)
clock := time.Now().Truncate(time.Hour)
l := New(h.Pool).WithClock(func() time.Time { return clock })
const callers = 40
rule := Rule{Name: "concurrent.rule", Limit: 10, Window: time.Minute}
var wg sync.WaitGroup
var mu sync.Mutex
allowed := 0
for i := 0; i < callers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
d, err := l.Allow(context.Background(), rule, Subject("hot-subject"))
if err != nil {
return
}
if d.Allowed {
mu.Lock()
allowed++
mu.Unlock()
}
}()
}
wg.Wait()
// EXACTLY the limit. Not "about" — if increments were lost, more would have
// been allowed; if the statement were not atomic, the count would be wrong
// in either direction.
if allowed != rule.Limit {
t.Errorf("%d of %d concurrent callers allowed, want exactly %d",
allowed, callers, rule.Limit)
}
var count int
if err := h.Pool.QueryRow(context.Background(),
`SELECT count FROM rate_limits WHERE bucket LIKE $1`, rule.Name+":%").Scan(&count); err != nil {
t.Fatalf("read counter: %v", err)
}
if count != callers {
t.Errorf("counter = %d, want %d — increments were lost", count, callers)
}
}
/* ── Failure behaviour ──────────────────────────────────────────────────── */
// A limiter that cannot count must refuse by default. Failing open turns a
// database blip into an unmetered window on the endpoints most worth guarding.
func TestFailsClosedByDefault(t *testing.T) {
l := New(brokenQuerier{}).WithClock(time.Now)
d, err := l.Allow(context.Background(), testRule, Subject("x"))
if err == nil {
t.Fatal("expected an error from a broken database")
}
if d.Allowed {
t.Error("the limiter failed OPEN by default; it must fail closed")
}
if d.RetryAfter <= 0 {
t.Error("a fail-closed decision carries no Retry-After")
}
}
func TestFailOpenIsOptIn(t *testing.T) {
l := New(brokenQuerier{}).WithFailOpen(true)
d, err := l.Allow(context.Background(), testRule, Subject("x"))
if err == nil {
t.Fatal("expected an error")
}
if !d.Allowed {
t.Error("WithFailOpen(true) did not permit the request")
}
}
/* ── Sweep ──────────────────────────────────────────────────────────────── */
func TestSweepRemovesOnlyExpiredWindows(t *testing.T) {
l, h, clock := newLimiter(t)
ctx := context.Background()
mustAllow(t, l, "old")
// Move past the old window, and open a new one.
*clock = clock.Add(2 * testRule.Window)
l.WithClock(func() time.Time { return *clock })
mustAllow(t, l, "current")
removed, err := l.Sweep(ctx, 100)
if err != nil {
t.Fatalf("Sweep: %v", err)
}
if removed != 1 {
t.Errorf("swept %d rows, want 1", removed)
}
// The live window must survive.
var remaining int
_ = h.Pool.QueryRow(ctx, `SELECT count(*) FROM rate_limits`).Scan(&remaining)
if remaining != 1 {
t.Errorf("%d rows left, want 1 — the live window was swept", remaining)
}
// Idempotent: a second sweep removes nothing and does not error.
if again, err := l.Sweep(ctx, 100); err != nil || again != 0 {
t.Errorf("second sweep: removed %d, err %v; want 0, nil", again, err)
}
}
func TestSweepIsBounded(t *testing.T) {
l, h, clock := newLimiter(t)
ctx := context.Background()
for i := 0; i < 10; i++ {
mustAllow(t, l, "subject-"+strings.Repeat("x", i))
}
*clock = clock.Add(2 * testRule.Window)
l.WithClock(func() time.Time { return *clock })
removed, err := l.Sweep(ctx, 4)
if err != nil {
t.Fatalf("Sweep: %v", err)
}
if removed != 4 {
t.Errorf("swept %d, want exactly the batch size 4", removed)
}
var left int
_ = h.Pool.QueryRow(ctx, `SELECT count(*) FROM rate_limits`).Scan(&left)
if left != 6 {
t.Errorf("%d rows left, want 6", left)
}
}
/* ── Helpers ────────────────────────────────────────────────────────────── */
// brokenQuerier fails every call, standing in for an unreachable database.
//
// Satisfies repo.Querier with pgx's real types, so this is the same interface
// the production limiter takes — a hand-rolled stand-in would prove the code
// works against a stand-in.
type brokenQuerier struct{}
func (brokenQuerier) Query(context.Context, string, ...any) (pgx.Rows, error) {
return nil, errBroken
}
func (brokenQuerier) QueryRow(context.Context, string, ...any) pgx.Row { return brokenRow{} }
func (brokenQuerier) Exec(context.Context, string, ...any) (pgconn.CommandTag, error) {
return pgconn.CommandTag{}, errBroken
}
type brokenRow struct{}
func (brokenRow) Scan(...any) error { return errBroken }
var errBroken = errString("ratelimit test: database unavailable")
type errString string
func (e errString) Error() string { return string(e) }
/* ── The fixed-window boundary, measured ────────────────────────────────── */
// The known limitation, asserted rather than assumed.
//
// A fixed window admits up to 2× the limit across a boundary: the limit at the
// end of one window and the limit again at the start of the next. This test
// measures that burst exactly, so the number in the documentation is a fact
// rather than a claim, and so a future change to the algorithm has to
// deliberately update it.
func TestFixedWindowBoundaryBurstIsExactlyTwice(t *testing.T) {
h := testutil.New(t)
// Start just inside a window, so "end of window" is reachable.
clock := time.Now().Truncate(time.Minute).Add(59 * time.Second)
l := New(h.Pool).WithClock(func() time.Time { return clock })
rule := Rule{Name: "boundary.rule", Limit: 5, Window: time.Minute}
allowed := 0
// Spend the whole limit at the very end of window 1.
for i := 0; i < rule.Limit; i++ {
d, err := l.Allow(context.Background(), rule, Subject("boundary"))
if err != nil {
t.Fatalf("Allow: %v", err)
}
if d.Allowed {
allowed++
}
}
// One second later, window 2 begins.
clock = clock.Add(time.Second)
l.WithClock(func() time.Time { return clock })
for i := 0; i < rule.Limit; i++ {
d, err := l.Allow(context.Background(), rule, Subject("boundary"))
if err != nil {
t.Fatalf("Allow: %v", err)
}
if d.Allowed {
allowed++
}
}
// Exactly 2× the limit in just over a second. This is the documented
// worst case — not worse, and not better.
if allowed != rule.Limit*2 {
t.Errorf("%d requests allowed across the boundary, want exactly %d (2× the limit)",
allowed, rule.Limit*2)
}
// And the burst does NOT continue: window 2's budget is now spent.
d, err := l.Allow(context.Background(), rule, Subject("boundary"))
if err != nil {
t.Fatalf("Allow: %v", err)
}
if d.Allowed {
t.Error("the burst continued past 2× the limit")
}
}
// Retry-After must point at the END of the current window, not at a fixed
// duration. A client told to wait a whole window when the window is nearly over
// waits twice as long as it needs to; one told to wait too little retries into
// the same refusal.
func TestRetryAfterPointsAtTheWindowBoundary(t *testing.T) {
h := testutil.New(t)
rule := Rule{Name: "retry.rule", Limit: 1, Window: time.Minute}
for _, offset := range []time.Duration{0, 15 * time.Second, 45 * time.Second, 59 * time.Second} {
clock := time.Now().Truncate(time.Minute).Add(offset)
l := New(h.Pool).WithClock(func() time.Time { return clock })
subject := Subject("retry-" + offset.String())
// Spend the budget, then be refused.
_, _ = l.Allow(context.Background(), rule, subject)
d, err := l.Allow(context.Background(), rule, subject)
if err != nil {
t.Fatalf("Allow: %v", err)
}
if d.Allowed {
t.Fatalf("offset %v: expected a refusal", offset)
}
want := rule.Window - offset
if d.RetryAfter != want {
t.Errorf("offset %v: RetryAfter = %v, want %v (the remainder of the window)",
offset, d.RetryAfter, want)
}
}
}