210 lines
7.4 KiB
Go
210 lines
7.4 KiB
Go
// Package ratelimit is a fixed-window request limiter shared across API
|
||
// instances.
|
||
//
|
||
// It exists because the limiter this service already had — httpserver's
|
||
// attemptLimiter — is an in-process map, and an in-process limiter behind N
|
||
// replicas enforces N times the configured limit. That is fine for the thing it
|
||
// guards (failed logins, where the real defence is the password hash's cost)
|
||
// and not fine for an endpoint that writes a database row for any caller who
|
||
// can reach it.
|
||
//
|
||
// The existing limiter is deliberately left alone. Replacing it is not this
|
||
// phase's job, it would change login behaviour, and the two have different
|
||
// shapes: attemptLimiter counts FAILURES and resets on success, which is the
|
||
// right model for a password and the wrong one for a request budget.
|
||
//
|
||
// WHAT THIS GUARANTEES
|
||
//
|
||
// - Correct under concurrency, including across instances: the count comes
|
||
// back from the same statement that increments it, so two callers cannot
|
||
// both read "9" and both proceed.
|
||
// - Bounded memory: the state is a table, swept by the cleanup job.
|
||
// - No credential ever becomes a key: callers pass an already-hashed subject,
|
||
// and Key refuses to build a bucket from anything that looks raw.
|
||
//
|
||
// # WHAT IT DOES NOT GUARANTEE
|
||
//
|
||
// Exactness at a window boundary. A fixed window admits up to 2× the limit
|
||
// across the seam — ten requests at 11:59:59 and ten more at 12:00:01. A
|
||
// sliding window would fix that and costs a row per request. For abuse
|
||
// prevention the burst is acceptable and the trade is deliberate.
|
||
package ratelimit
|
||
|
||
import (
|
||
"context"
|
||
"crypto/sha256"
|
||
"encoding/hex"
|
||
"fmt"
|
||
"time"
|
||
|
||
"github.com/krow/krow-backend/go-api/internal/repo"
|
||
)
|
||
|
||
// Limiter counts requests per bucket per window.
|
||
type Limiter struct {
|
||
db repo.Querier
|
||
now func() time.Time
|
||
|
||
// failOpen decides what happens when the DATABASE fails, which is the one
|
||
// judgement call in this package.
|
||
//
|
||
// Default false — fail CLOSED. A limiter that cannot count is a limiter
|
||
// that is not limiting, and these endpoints are the ones worth protecting
|
||
// most when things are already going wrong. The alternative, failing open,
|
||
// turns a database blip into an unmetered window on an endpoint that
|
||
// writes rows for anonymous callers.
|
||
//
|
||
// Configurable because that is not the right answer everywhere: a
|
||
// deployment that would rather serve MCP degraded than refuse it can say
|
||
// so deliberately, in one place, rather than by a comment somebody has to
|
||
// remember.
|
||
failOpen bool
|
||
}
|
||
|
||
// New builds a limiter over the shared pool.
|
||
func New(db repo.Querier) *Limiter {
|
||
return &Limiter{db: db, now: time.Now}
|
||
}
|
||
|
||
// WithClock replaces the clock, so window rollover can be tested without
|
||
// waiting for one.
|
||
func (l *Limiter) WithClock(now func() time.Time) *Limiter {
|
||
l.now = now
|
||
return l
|
||
}
|
||
|
||
// WithFailOpen makes a database failure permit the request rather than refuse
|
||
// it. See the field comment: the default is to refuse.
|
||
func (l *Limiter) WithFailOpen(open bool) *Limiter {
|
||
l.failOpen = open
|
||
return l
|
||
}
|
||
|
||
// Rule is one limit: how many requests, over how long.
|
||
type Rule struct {
|
||
// Name is the scope, and it becomes the bucket's prefix. Keep it stable —
|
||
// renaming a scope resets everyone's counter.
|
||
Name string
|
||
|
||
// Limit is the number of requests permitted per window.
|
||
Limit int
|
||
|
||
// Window is the fixed window's length.
|
||
Window time.Duration
|
||
}
|
||
|
||
// Decision is the answer for one request.
|
||
type Decision struct {
|
||
// Allowed is whether the caller may proceed.
|
||
Allowed bool
|
||
|
||
// Remaining is how many requests are left in this window, never negative.
|
||
Remaining int
|
||
|
||
// RetryAfter is how long until the window rolls over. Rendered into the
|
||
// Retry-After header on a 429, so a well-behaved client waits exactly long
|
||
// enough rather than guessing.
|
||
RetryAfter time.Duration
|
||
|
||
// Limit and Window echo the rule, for the response headers.
|
||
Limit int
|
||
Window time.Duration
|
||
}
|
||
|
||
// Subject hashes a bucket subject.
|
||
//
|
||
// EVERY caller must pass identifying material through this. A bucket key built
|
||
// from a raw token would write that token to a table, to any log line naming
|
||
// the bucket, and to every slow-query report the row ever appears in. Hashing
|
||
// costs nothing here — the value is never read back, only compared.
|
||
//
|
||
// Truncated to 32 hex characters: 128 bits, far beyond collision risk for a
|
||
// counter, and it keeps the keys readable in a psql session while still being
|
||
// irreversible.
|
||
func Subject(raw string) string {
|
||
sum := sha256.Sum256([]byte(raw))
|
||
return hex.EncodeToString(sum[:])[:32]
|
||
}
|
||
|
||
// Allow records one request against a rule and reports whether it may proceed.
|
||
//
|
||
// The whole decision is one statement. It is worth reading, because everything
|
||
// this package claims about concurrency rests on it:
|
||
//
|
||
// INSERT INTO rate_limits (bucket, window_start, count, expires_at)
|
||
// VALUES ($1, $2, 1, $3)
|
||
// ON CONFLICT (bucket, window_start)
|
||
// DO UPDATE SET count = rate_limits.count + 1
|
||
// RETURNING count
|
||
//
|
||
// The row is created or incremented, and the resulting count comes back, in one
|
||
// round trip under one implicit transaction. Two instances racing on the same
|
||
// bucket serialise on the primary key, and each sees a distinct count. There is
|
||
// no read-then-write window for them to slip through.
|
||
//
|
||
// A request is counted even when it is refused. That is deliberate: a caller
|
||
// hammering a limit should not be able to keep their own window open by
|
||
// spending it, and the alternative — not counting refusals — makes the limit
|
||
// cheaper to probe.
|
||
func (l *Limiter) Allow(ctx context.Context, rule Rule, subject string) (Decision, error) {
|
||
now := l.now()
|
||
windowStart := now.Truncate(rule.Window)
|
||
expiresAt := windowStart.Add(rule.Window)
|
||
bucket := rule.Name + ":" + subject
|
||
|
||
var count int
|
||
err := l.db.QueryRow(ctx,
|
||
`INSERT INTO rate_limits (bucket, window_start, count, expires_at)
|
||
VALUES ($1, $2, 1, $3)
|
||
ON CONFLICT (bucket, window_start)
|
||
DO UPDATE SET count = rate_limits.count + 1
|
||
RETURNING count`,
|
||
bucket, windowStart, expiresAt).Scan(&count)
|
||
if err != nil {
|
||
if l.failOpen {
|
||
return Decision{Allowed: true, Remaining: rule.Limit, Limit: rule.Limit, Window: rule.Window},
|
||
fmt.Errorf("ratelimit: %w", err)
|
||
}
|
||
return Decision{Allowed: false, RetryAfter: rule.Window, Limit: rule.Limit, Window: rule.Window},
|
||
fmt.Errorf("ratelimit: %w", err)
|
||
}
|
||
|
||
remaining := rule.Limit - count
|
||
if remaining < 0 {
|
||
remaining = 0
|
||
}
|
||
return Decision{
|
||
Allowed: count <= rule.Limit,
|
||
Remaining: remaining,
|
||
RetryAfter: expiresAt.Sub(now),
|
||
Limit: rule.Limit,
|
||
Window: rule.Window,
|
||
}, nil
|
||
}
|
||
|
||
// Sweep deletes expired counters, in bounded batches.
|
||
//
|
||
// Bounded because an unbounded DELETE on a busy table takes a lock for as long
|
||
// as it takes to finish, and "as long as it takes" is not a number anyone can
|
||
// predict at 3am. A batch of a few thousand rows completes in milliseconds and
|
||
// can simply be run again.
|
||
//
|
||
// Safe to run concurrently: two sweeps delete disjoint sets because the
|
||
// subquery re-reads under each statement's own snapshot, and a row deleted
|
||
// twice is not an error.
|
||
func (l *Limiter) Sweep(ctx context.Context, batch int) (int64, error) {
|
||
if batch <= 0 {
|
||
batch = 5000
|
||
}
|
||
tag, err := l.db.Exec(ctx,
|
||
`DELETE FROM rate_limits
|
||
WHERE ctid IN (
|
||
SELECT ctid FROM rate_limits WHERE expires_at < $1 LIMIT $2
|
||
)`,
|
||
l.now(), batch)
|
||
if err != nil {
|
||
return 0, fmt.Errorf("ratelimit: sweep: %w", err)
|
||
}
|
||
return tag.RowsAffected(), nil
|
||
}
|