453 lines
14 KiB
Go
453 lines
14 KiB
Go
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)
|
||
}
|
||
}
|
||
}
|