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) } } }