From 1aa524fb22de900491dcaa6b315890b98bda99ae Mon Sep 17 00:00:00 2001 From: DcruBro Date: Sun, 6 Sep 2026 15:15:21 +0200 Subject: [PATCH] starter code --- internal/api/auth.go | 322 +++++++++++++++++++++++++++ internal/api/auth_test.go | 178 +++++++++++++++ internal/api/errors.go | 57 +++++ internal/api/store.go | 47 ++++ internal/api/whoami.go | 22 ++ internal/ratelimit/memory.go | 80 +++++++ internal/ratelimit/ratelimit.go | 62 ++++++ internal/ratelimit/ratelimit_test.go | 208 +++++++++++++++++ internal/ratelimit/redis.go | 98 ++++++++ internal/ratelimit/redis_adapter.go | 20 ++ 10 files changed, 1094 insertions(+) create mode 100644 internal/api/auth.go create mode 100644 internal/api/auth_test.go create mode 100644 internal/api/errors.go create mode 100644 internal/api/store.go create mode 100644 internal/api/whoami.go create mode 100644 internal/ratelimit/memory.go create mode 100644 internal/ratelimit/ratelimit.go create mode 100644 internal/ratelimit/ratelimit_test.go create mode 100644 internal/ratelimit/redis.go create mode 100644 internal/ratelimit/redis_adapter.go diff --git a/internal/api/auth.go b/internal/api/auth.go new file mode 100644 index 0000000..0567d01 --- /dev/null +++ b/internal/api/auth.go @@ -0,0 +1,322 @@ +package api + +import ( + "context" + "errors" + "log/slog" + "net/http" + "strconv" + "strings" + "time" + + "git.dcrubro.com/dcrubro/hammy-backend/internal/apikey" + "git.dcrubro.com/dcrubro/hammy-backend/internal/ratelimit" +) + +// Scope names. These must match the keys_scopes_known CHECK in +// migrations/003_api.sql; the database is the source of truth and will reject +// anything not listed there. +const ( + ScopeReference = "reference" + ScopeCallsign = "callsign" + ScopeSpots = "spots" + ScopePropagation = "propagation" + ScopeLogbook = "logbook" + ScopeNet44 = "net44" + ScopeAdmin = "admin" +) + +// ErrKeyNotFound is what a KeyStore returns when no active key matches. +var ErrKeyNotFound = errors.New("api: key not found") + +// KeyRecord is what authentication needs from the database. Deliberately not +// the sqlc-generated type. +type KeyRecord struct { + ID int64 + OwnerID int64 + Scopes []string + QuotaTier string + Status string + OwnerStatus string + ExpiresAt *time.Time +} + +// KeyStore is the database slice this middleware needs. +type KeyStore interface { + KeyByHash(ctx context.Context, hash []byte) (KeyRecord, error) +} + +// Principal is the authenticated caller, attached to the request context. +type Principal struct { + KeyID int64 + OwnerID int64 + Scopes []string + QuotaTier string + + // EndUser is the caller's own identifier for whoever triggered the + // request - for the bot, a hashed Discord user ID. Optional, but without + // it a single user can consume a whole instance's quota, and the + // per-user counter is the enumeration signal worth watching. + EndUser string +} + +// Has reports whether the principal holds a scope. Admin is NOT a wildcard - +// making it one means a single compromised admin key reaches everything, and +// makes it impossible to tell from the database what a key can actually do. +func (p *Principal) Has(scope string) bool { + if p == nil { + return false + } + + for _, s := range p.Scopes { + if s == scope { + return true + } + } + + return false +} + +type contextKey struct{} + +var principalKey contextKey + +// PrincipalFrom returns the authenticated caller, or nil if the request did not +// pass through Authenticate. +func PrincipalFrom(ctx context.Context) *Principal { + p, _ := ctx.Value(principalKey).(*Principal) + + return p +} + +// Authenticator resolves a bearer token to a Principal and enforces quota. +type Authenticator struct { + Keys KeyStore + Limiter ratelimit.Limiter + Logger *slog.Logger + + // Now is injectable so expiry can be tested without waiting. + Now func() time.Time + + // EndUserHeader carries the caller's identifier for the person on whose + // behalf the request is made. Empty disables per-user limiting. + EndUserHeader string +} + +func (a *Authenticator) now() time.Time { + if a.Now != nil { + return a.Now() + } + + return time.Now() +} + +func (a *Authenticator) log() *slog.Logger { + if a.Logger != nil { + return a.Logger + } + + return slog.Default() +} + +// bearer pulls the token out of an Authorization header. +func bearer(r *http.Request) string { + h := r.Header.Get("Authorization") + if h == "" { + return "" + } + + // Scheme is case-insensitive per RFC 7235, and clients get this wrong. + const prefix = "bearer " + if len(h) < len(prefix) || !strings.EqualFold(h[:len(prefix)], prefix) { + return "" + } + + return strings.TrimSpace(h[len(prefix):]) +} + +// Authenticate is the middleware. Every route that touches user data or an +// upstream feed must sit behind it. +// +// The bot is NOT a trusted client: it is GPL, so anyone can fork it, strip its +// checks and point it at this API. Scope and quota decisions therefore live +// here, on the server, and the bot's versions of them are a UX nice-ty that +// saves a round trip. +func (a *Authenticator) Authenticate(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + token := bearer(r) + if token == "" { + WriteError(w, http.StatusUnauthorized, CodeMissingKey, + "Provide an API key as 'Authorization: Bearer hmy_live_...'.") + + return + } + + // Structure and checksum first: a flood of garbage headers costs no + // database queries. + if _, err := apikey.Validate(token); err != nil { + WriteError(w, http.StatusUnauthorized, CodeInvalidKey, + "That does not look like a Hammy API key.") + + return + } + + rec, err := a.Keys.KeyByHash(r.Context(), apikey.Hash(token)) + if err != nil { + if errors.Is(err, ErrKeyNotFound) { + WriteError(w, http.StatusUnauthorized, CodeInvalidKey, + "Unknown API key.") + + return + } + + a.log().Error("key lookup failed", "err", err, "key", apikey.Redact(token)) + WriteError(w, http.StatusInternalServerError, CodeInternal, + "Could not verify that key. Try again shortly.") + + return + } + + if code, msg, ok := checkUsable(&rec, a.now()); !ok { + status := http.StatusUnauthorized + if code == CodeOwnerBlocked { + status = http.StatusForbidden + } + + WriteError(w, status, code, msg) + + return + } + + principal := &Principal{ + KeyID: rec.ID, + OwnerID: rec.OwnerID, + Scopes: rec.Scopes, + QuotaTier: rec.QuotaTier, + } + + if a.EndUserHeader != "" { + principal.EndUser = r.Header.Get(a.EndUserHeader) + } + + if !a.allow(w, r, principal) { + return + } + + ctx := context.WithValue(r.Context(), principalKey, principal) + next.ServeHTTP(w, r.WithContext(ctx)) + }) +} + +// checkUsable applies the status and expiry rules. Split out so it can be +// tested directly and so the reasons stay in one readable place. +func checkUsable(rec *KeyRecord, now time.Time) (code, message string, ok bool) { + switch rec.Status { + case "revoked": + return CodeKeyRevoked, "This key has been revoked.", false + case "suspended": + return CodeKeySuspended, "This key is suspended. Check your email for the reason.", false + case "active": + // fall through + default: + return CodeInvalidKey, "This key is not usable.", false + } + + if rec.OwnerStatus != "active" { + return CodeOwnerBlocked, "This account is not active.", false + } + + if rec.ExpiresAt != nil && !rec.ExpiresAt.After(now) { + return CodeKeyExpired, "This key expired. Issue a new one and rotate.", false + } + + return "", "", true +} + +// allow enforces both ceilings and writes the 429 itself. Returns false if the +// request must stop. +func (a *Authenticator) allow(w http.ResponseWriter, r *http.Request, p *Principal) bool { + if a.Limiter == nil { + return true + } + + quota := ratelimit.QuotaFor(p.QuotaTier) + + res, err := a.Limiter.Allow(r.Context(), + "key:"+strconv.FormatInt(p.KeyID, 10), quota.KeyLimit, quota.Window) + if err != nil { + // Fail OPEN. A Redis outage taking the whole API down with it is a + // worse outcome than briefly unmetered traffic - the limiter protects + // against abuse, it is not an authentication control. + a.log().Error("rate limiter unavailable, allowing", "err", err, "key_id", p.KeyID) + + return true + } + + writeLimitHeaders(w, res) + + if !res.Allowed { + WriteError(w, http.StatusTooManyRequests, CodeRateLimited, + "Rate limit exceeded for this API key.") + + return false + } + + // Per-user ceiling, if the caller told us who this is for. + if p.EndUser != "" && quota.UserLimit > 0 { + userRes, err := a.Limiter.Allow(r.Context(), + "user:"+strconv.FormatInt(p.KeyID, 10)+":"+p.EndUser, + quota.UserLimit, quota.Window) + if err != nil { + a.log().Error("per-user rate limiter unavailable, allowing", "err", err) + + return true + } + + if !userRes.Allowed { + writeLimitHeaders(w, userRes) + WriteError(w, http.StatusTooManyRequests, CodeRateLimited, + "Rate limit exceeded for this user.") + + return false + } + } + + return true +} + +func writeLimitHeaders(w http.ResponseWriter, res ratelimit.Result) { + if res.Limit <= 0 { + return + } + + w.Header().Set("X-RateLimit-Limit", strconv.Itoa(res.Limit)) + w.Header().Set("X-RateLimit-Remaining", strconv.Itoa(res.Remaining)) + + if !res.Allowed { + secs := int(res.RetryAfter.Seconds()) + if secs < 1 { + secs = 1 + } + + w.Header().Set("Retry-After", strconv.Itoa(secs)) + } +} + +// RequireScope wraps a handler so it only runs for principals holding the +// scope. Authorisation is separate from authentication on purpose - a key being +// valid says nothing about what it may do. +func RequireScope(scope string, next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + p := PrincipalFrom(r.Context()) + + if !p.Has(scope) { + WriteError(w, http.StatusForbidden, CodeForbidden, + "This key does not have the '"+scope+"' scope.") + + return + } + + next.ServeHTTP(w, r) + }) +} diff --git a/internal/api/auth_test.go b/internal/api/auth_test.go new file mode 100644 index 0000000..057b8b7 --- /dev/null +++ b/internal/api/auth_test.go @@ -0,0 +1,178 @@ +package api + +import ( + "bytes" + "errors" + "strings" + "testing" +) + +func TestGenerateRoundTrip(t *testing.T) { + for _, env := range []Environment{Live, Test} { + g, err := Generate(env) + if err != nil { + t.Fatalf("Generate(%s): %v", env, err) + } + + got, err := Validate(g.Plaintext) + if err != nil { + t.Fatalf("Validate(%q): %v", g.Plaintext, err) + } + + if got != env { + t.Errorf("environment round trip: got %q want %q", got, env) + } + + if !strings.HasPrefix(g.Plaintext, Issuer+"_"+string(env)+"_") { + t.Errorf("prefix wrong: %q", g.Plaintext) + } + + if len(g.Hash) != 32 { + t.Errorf("hash length %d, want 32 (the api.keys CHECK requires it)", len(g.Hash)) + } + + if !bytes.Equal(g.Hash, Hash(g.Plaintext)) { + t.Error("Generate's hash differs from Hash of its own plaintext") + } + + if len(g.Display) != displayLen { + t.Errorf("display length %d, want %d", len(g.Display), displayLen) + } + + if strings.Contains(g.Plaintext[len(g.Display):], "") && g.Display == g.Plaintext { + t.Error("display prefix is the whole key") + } + } +} + +func TestGenerateIsUnique(t *testing.T) { + const n = 2000 + + seen := make(map[string]bool, n) + + for i := 0; i < n; i++ { + g, err := Generate(Live) + if err != nil { + t.Fatalf("Generate: %v", err) + } + + if seen[g.Plaintext] { + t.Fatalf("duplicate key after %d generations", i) + } + + seen[g.Plaintext] = true + } +} + +func TestGenerateRejectsUnknownEnvironment(t *testing.T) { + if _, err := Generate(Environment("prod")); !errors.Is(err, ErrEnvironment) { + t.Errorf("got %v, want ErrEnvironment", err) + } +} + +func TestValidateRejectsMalformed(t *testing.T) { + good, err := Generate(Live) + if err != nil { + t.Fatal(err) + } + + cases := []struct { + name string + key string + want error + }{ + {"empty", "", ErrMalformed}, + {"no separators", "notakey", ErrMalformed}, + {"one separator", "hmy_live", ErrMalformed}, + {"wrong issuer", "xxx_live_" + good.Plaintext[9:], ErrMalformed}, + {"unknown env", "hmy_prod_" + good.Plaintext[9:], ErrEnvironment}, + {"body too short", "hmy_live_abc", ErrMalformed}, + {"body too long", good.Plaintext + "X", ErrMalformed}, + {"body truncated by one", good.Plaintext[:len(good.Plaintext)-1], ErrMalformed}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + if _, err := Validate(c.key); !errors.Is(err, c.want) { + t.Errorf("Validate(%q) = %v, want %v", c.key, err, c.want) + } + }) + } +} + +// The point of the checksum: a key with a typo is rejected without a database +// round trip. +func TestValidateCatchesTampering(t *testing.T) { + g, err := Generate(Live) + if err != nil { + t.Fatal(err) + } + + body := []byte(g.Plaintext) + caught := 0 + + // Flip one character at each position in the secret and confirm the + // checksum notices. + for i := len("hmy_live_"); i < len(body)-checksumLen; i++ { + orig := body[i] + + if orig == 'A' { + body[i] = 'B' + } else { + body[i] = 'A' + } + + if _, err := Validate(string(body)); errors.Is(err, ErrBadChecksum) { + caught++ + } + + body[i] = orig + } + + total := len(body) - checksumLen - len("hmy_live_") + + if caught != total { + t.Errorf("checksum caught %d/%d single-character corruptions", caught, total) + } +} + +func TestHashIsStable(t *testing.T) { + const key = "hmy_live_0123456789abcdefghijklmnopqrstuvwxyzABCDEFGabcdef" + + a, b := Hash(key), Hash(key) + + if !bytes.Equal(a, b) { + t.Error("Hash is not deterministic") + } + + if bytes.Equal(a, Hash(key+"x")) { + t.Error("Hash collided on different input") + } +} + +func TestRedact(t *testing.T) { + g, err := Generate(Live) + if err != nil { + t.Fatal(err) + } + + r := Redact(g.Plaintext) + + if strings.Contains(g.Plaintext, r) && len(r) >= len(g.Plaintext) { + t.Error("Redact returned the whole key") + } + + if len(r) > displayLen+3 { + t.Errorf("Redact returned %d chars, too much", len(r)) + } + + // A secret must never survive redaction. + secret := g.Plaintext[len("hmy_live_"):] + if strings.Contains(r, secret) { + t.Error("Redact leaked the secret") + } + + if got := Redact("short"); got != "hmy_***" { + t.Errorf("Redact(short) = %q", got) + } +} diff --git a/internal/api/errors.go b/internal/api/errors.go new file mode 100644 index 0000000..0960f88 --- /dev/null +++ b/internal/api/errors.go @@ -0,0 +1,57 @@ +package api + +import ( + "encoding/json" + "log/slog" + "net/http" +) + +// Error is the single shape every failure takes. Clients get a stable machine +// readable code plus a human message. A self-hosted bot instance can branch on +// the code without parsing prose. +type Error struct { + Code string `json:"code"` + Message string `json:"message"` +} + +type errorEnvelope struct { + Error Error `json:"error"` +} + +const ( + CodeMissingKey = "missing_key" + CodeInvalidKey = "invalid_key" + CodeKeyRevoked = "key_revoked" + CodeKeySuspended = "key_suspended" + CodeKeyExpired = "key_expired" + CodeOwnerBlocked = "owner_blocked" + CodeForbidden = "insufficient_scope" + CodeRateLimited = "rate_limited" + CodeInternal = "internal_error" +) + +// WriteError sends a JSON error. Deliberately specific about WHY authentication +// failed rather than a uniform "unauthorized". +// +// The usual argument for being vague is to prevent enumeration, but a key is +// 256 bits from crypto/rand - there is nothing to enumerate. Telling a +// developer their key expired, rather than making them guess, is worth far more +// than the non-existent secrecy gained. +func WriteError(w http.ResponseWriter, status int, code, message string) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(status) + + if err := json.NewEncoder(w).Encode(errorEnvelope{Error{Code: code, Message: message}}); err != nil { + slog.Error("writing error response", "err", err) + } +} + +// WriteJSON sends a success payload. +func WriteJSON(w http.ResponseWriter, status int, v any) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(status) + + if err := json.NewEncoder(w).Encode(v); err != nil { + slog.Error("writing response", "err", err) + } +} diff --git a/internal/api/store.go b/internal/api/store.go new file mode 100644 index 0000000..d144127 --- /dev/null +++ b/internal/api/store.go @@ -0,0 +1,47 @@ +package api + +import ( + "context" + "errors" + + "github.com/jackc/pgx/v5" + + "git.dcrubro.com/dcrubro/hammy-backend/internal/db" +) + +// PgKeyStore adapts the sqlc-generated queries to the KeyStore interface. +// The indirection earns its keep twice: the middleware is testable without a +// database, and adding an unrelated column to api.keys does not ripple into the +// auth package. +type PgKeyStore struct { + Q *db.Queries +} + +func (s PgKeyStore) KeyByHash(ctx context.Context, hash []byte) (KeyRecord, error) { + row, err := s.Q.KeyByHash(ctx, hash) + if err != nil { + // pgx returns ErrNoRows for an empty :one result. Translate it so the + // middleware never has to know which driver is underneath. + if errors.Is(err, pgx.ErrNoRows) { + return KeyRecord{}, ErrKeyNotFound + } + + return KeyRecord{}, err + } + + rec := KeyRecord{ + ID: row.ID, + OwnerID: row.OwnerID, + Scopes: row.Scopes, + QuotaTier: row.QuotaTier, + Status: row.Status, + OwnerStatus: row.OwnerStatus, + } + + // sqlc with emit_pointers_for_null_types gives *time.Time for a nullable + // timestamptz. If your generated code uses pgtype.Timestamptz instead, + // convert here rather than in the middleware. + rec.ExpiresAt = row.ExpiresAt + + return rec, nil +} diff --git a/internal/api/whoami.go b/internal/api/whoami.go new file mode 100644 index 0000000..f16dedf --- /dev/null +++ b/internal/api/whoami.go @@ -0,0 +1,22 @@ +package api + +import ( + "net/http" +) + +// Whoami echoes the resolved principal. The cheapest possible end-to-end test +// of the auth chain: if this returns your key id and scopes, then the bearer +// header, the checksum, the hash lookup, the status checks and the rate limiter +// are all working. +func Whoami() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + p := PrincipalFrom(r.Context()) + + WriteJSON(w, http.StatusOK, map[string]any{ + "key_id": p.KeyID, + "owner_id": p.OwnerID, + "scopes": p.Scopes, + "quota_tier": p.QuotaTier, + }) + }) +} diff --git a/internal/ratelimit/memory.go b/internal/ratelimit/memory.go new file mode 100644 index 0000000..d8e2909 --- /dev/null +++ b/internal/ratelimit/memory.go @@ -0,0 +1,80 @@ +package ratelimit + +import ( + "context" + "sync" + "time" +) + +// Memory is a fixed-window limiter held in process. Useful for tests and for a +// single-instance self-hoster who would rather not run Redis. +// +// Not suitable for more than one backend process: each would keep its own +// counters, so N processes allow N times the limit. +type Memory struct { + mu sync.Mutex + windows map[string]*window + Now func() time.Time // injectable so tests need not sleep +} + +type window struct { + count int + expires time.Time +} + +func NewMemory() *Memory { + return &Memory{ + windows: make(map[string]*window), + Now: time.Now, + } +} + +func (m *Memory) Allow(_ context.Context, key string, limit int, dur time.Duration) (Result, error) { + if limit <= 0 { + return Result{Allowed: true, Limit: 0, Remaining: 0}, nil + } + + m.mu.Lock() + defer m.mu.Unlock() + + now := m.Now() + + w, ok := m.windows[key] + if !ok || now.After(w.expires) { + w = &window{expires: now.Add(dur)} + m.windows[key] = w + } + + w.count++ + + remaining := limit - w.count + if remaining < 0 { + remaining = 0 + } + + return Result{ + Allowed: w.count <= limit, + Limit: limit, + Remaining: remaining, + RetryAfter: w.expires.Sub(now), + }, nil +} + +// Sweep drops expired windows. Without it the map grows for every key ever +// seen, which for per-user keys is unbounded. Call it periodically. +func (m *Memory) Sweep() int { + m.mu.Lock() + defer m.mu.Unlock() + + now := m.Now() + removed := 0 + + for k, w := range m.windows { + if now.After(w.expires) { + delete(m.windows, k) + removed++ + } + } + + return removed +} diff --git a/internal/ratelimit/ratelimit.go b/internal/ratelimit/ratelimit.go new file mode 100644 index 0000000..6adf95c --- /dev/null +++ b/internal/ratelimit/ratelimit.go @@ -0,0 +1,62 @@ +// Package ratelimit enforces per-key and per-end-user request ceilings. +// +// Redis is authoritative for the live window; api.usage_daily holds completed +// days. A Redis flush therefore loses at most the window in progress, which for +// abuse detection is an acceptable trade. +package ratelimit + +import ( + "context" + "time" +) + +// Result is what a limiter says about one request. +type Result struct { + Allowed bool + Limit int + Remaining int + RetryAfter time.Duration +} + +// Limiter is deliberately narrow so the middleware can be tested without Redis, +// and so a self-hoster running a single instance could drop in an in-process +// implementation instead of standing up Redis at all. +type Limiter interface { + // Allow records one request against key and reports whether it may + // proceed. Implementations must be atomic: a check-then-increment race + // under load is exactly when the limit matters most. + Allow(ctx context.Context, key string, limit int, window time.Duration) (Result, error) +} + +// Quota is the ceiling for one tier, from api.keys.quota_tier. +type Quota struct { + // Requests allowed per window for the key as a whole. + KeyLimit int + + // Requests allowed per window for a single end user within that key. One + // Discord user should not be able to consume a whole self-hosted + // instance's allowance, and a key whose traffic is 90% one user is the + // enumeration signal worth alerting on. + UserLimit int + + Window time.Duration +} + +// Tiers maps api.keys.quota_tier to its ceiling. Values are a starting point, +// not a considered policy: watch real usage before hardening them. +var Tiers = map[string]Quota{ + "default": {KeyLimit: 600, UserLimit: 120, Window: time.Minute}, + "verified": {KeyLimit: 3000, UserLimit: 600, Window: time.Minute}, + "trusted": {KeyLimit: 12000, UserLimit: 2400, Window: time.Minute}, + "internal": {KeyLimit: 0, UserLimit: 0, Window: time.Minute}, // 0 = unlimited +} + +// QuotaFor falls back to the tightest tier for an unknown name, so a typo in +// the database cannot accidentally grant unlimited access. +func QuotaFor(tier string) Quota { + if q, ok := Tiers[tier]; ok { + return q + } + + return Tiers["default"] +} diff --git a/internal/ratelimit/ratelimit_test.go b/internal/ratelimit/ratelimit_test.go new file mode 100644 index 0000000..639fb07 --- /dev/null +++ b/internal/ratelimit/ratelimit_test.go @@ -0,0 +1,208 @@ +package ratelimit + +import ( + "context" + "errors" + "testing" + "time" +) + +func TestMemoryWindow(t *testing.T) { + now := time.Date(2026, 9, 4, 12, 0, 0, 0, time.UTC) + m := NewMemory() + m.Now = func() time.Time { return now } + + ctx := context.Background() + + for i := 1; i <= 3; i++ { + res, err := m.Allow(ctx, "k", 3, time.Minute) + if err != nil || !res.Allowed { + t.Fatalf("request %d: allowed=%v err=%v", i, res.Allowed, err) + } + if res.Remaining != 3-i { + t.Errorf("request %d: remaining %d, want %d", i, res.Remaining, 3-i) + } + } + + res, _ := m.Allow(ctx, "k", 3, time.Minute) + if res.Allowed { + t.Error("4th request allowed past a limit of 3") + } + if res.RetryAfter <= 0 { + t.Error("RetryAfter not set on a rejection") + } + + // Window rolls over. + now = now.Add(time.Minute + time.Second) + if res, _ := m.Allow(ctx, "k", 3, time.Minute); !res.Allowed { + t.Error("request rejected after the window expired") + } +} + +func TestMemoryKeysAreIndependent(t *testing.T) { + m := NewMemory() + ctx := context.Background() + + for i := 0; i < 5; i++ { + m.Allow(ctx, "a", 5, time.Minute) + } + + if res, _ := m.Allow(ctx, "a", 5, time.Minute); res.Allowed { + t.Error("key a not limited") + } + if res, _ := m.Allow(ctx, "b", 5, time.Minute); !res.Allowed { + t.Error("key b affected by key a") + } +} + +func TestMemoryZeroLimitIsUnlimited(t *testing.T) { + m := NewMemory() + for i := 0; i < 1000; i++ { + if res, _ := m.Allow(context.Background(), "k", 0, time.Minute); !res.Allowed { + t.Fatalf("zero limit rejected at %d", i) + } + } +} + +func TestMemorySweep(t *testing.T) { + now := time.Now() + m := NewMemory() + m.Now = func() time.Time { return now } + + m.Allow(context.Background(), "a", 5, time.Minute) + m.Allow(context.Background(), "b", 5, time.Hour) + + if n := m.Sweep(); n != 0 { + t.Errorf("swept %d live windows", n) + } + + now = now.Add(2 * time.Minute) + + if n := m.Sweep(); n != 1 { + t.Errorf("swept %d, want 1 (a expired, b did not)", n) + } +} + +func TestMemoryIsConcurrencySafe(t *testing.T) { + m := NewMemory() + ctx := context.Background() + + const goroutines, each = 20, 50 + + done := make(chan int, goroutines) + + for g := 0; g < goroutines; g++ { + go func() { + allowed := 0 + for i := 0; i < each; i++ { + if res, _ := m.Allow(ctx, "shared", 100, time.Minute); res.Allowed { + allowed++ + } + } + done <- allowed + }() + } + + total := 0 + for g := 0; g < goroutines; g++ { + total += <-done + } + + // Exactly 100 of the 1000 attempts may pass. A check-then-increment race + // would let more through. + if total != 100 { + t.Errorf("%d requests allowed under concurrency, want exactly 100", total) + } +} + +// fakeRunner stands in for Redis so the script result handling can be tested. +type fakeRunner struct { + ret any + err error + keys []string + args []any +} + +func (f *fakeRunner) Eval(_ context.Context, _ string, keys []string, args ...any) (any, error) { + f.keys, f.args = keys, args + return f.ret, f.err +} + +func TestRedisResultParsing(t *testing.T) { + ctx := context.Background() + + t.Run("under limit", func(t *testing.T) { + f := &fakeRunner{ret: []any{int64(3), int64(45000)}} + r := NewRedis(f) + + res, err := r.Allow(ctx, "key:1", 10, time.Minute) + if err != nil { + t.Fatal(err) + } + if !res.Allowed || res.Remaining != 7 { + t.Errorf("allowed=%v remaining=%d, want true/7", res.Allowed, res.Remaining) + } + if res.RetryAfter != 45*time.Second { + t.Errorf("RetryAfter %v, want 45s", res.RetryAfter) + } + if len(f.keys) != 1 || f.keys[0] != "rl:key:1" { + t.Errorf("keys %v, want [rl:key:1]", f.keys) + } + }) + + t.Run("over limit", func(t *testing.T) { + r := NewRedis(&fakeRunner{ret: []any{int64(11), int64(1000)}}) + res, _ := r.Allow(ctx, "k", 10, time.Minute) + if res.Allowed || res.Remaining != 0 { + t.Errorf("allowed=%v remaining=%d, want false/0", res.Allowed, res.Remaining) + } + }) + + t.Run("exactly at limit is allowed", func(t *testing.T) { + r := NewRedis(&fakeRunner{ret: []any{int64(10), int64(1000)}}) + if res, _ := r.Allow(ctx, "k", 10, time.Minute); !res.Allowed { + t.Error("request at exactly the limit was rejected") + } + }) + + t.Run("plain int accepted", func(t *testing.T) { + r := NewRedis(&fakeRunner{ret: []any{1, 1000}}) + if _, err := r.Allow(ctx, "k", 10, time.Minute); err != nil { + t.Errorf("int result rejected: %v", err) + } + }) + + t.Run("missing ttl falls back to the window", func(t *testing.T) { + r := NewRedis(&fakeRunner{ret: []any{int64(11), int64(-1)}}) + res, _ := r.Allow(ctx, "k", 10, time.Minute) + if res.RetryAfter != time.Minute { + t.Errorf("RetryAfter %v, want 1m", res.RetryAfter) + } + }) + + t.Run("eval error propagates", func(t *testing.T) { + r := NewRedis(&fakeRunner{err: errors.New("nope")}) + if _, err := r.Allow(ctx, "k", 10, time.Minute); err == nil { + t.Error("expected an error so the caller can fail open") + } + }) + + t.Run("garbage result is an error not a panic", func(t *testing.T) { + r := NewRedis(&fakeRunner{ret: "not a list"}) + if _, err := r.Allow(ctx, "k", 10, time.Minute); err == nil { + t.Error("expected an error") + } + }) + + t.Run("zero limit skips redis entirely", func(t *testing.T) { + f := &fakeRunner{ret: []any{int64(1), int64(1)}} + r := NewRedis(f) + res, _ := r.Allow(ctx, "k", 0, time.Minute) + if !res.Allowed { + t.Error("zero limit rejected") + } + if f.keys != nil { + t.Error("redis was called for an unlimited tier") + } + }) +} diff --git a/internal/ratelimit/redis.go b/internal/ratelimit/redis.go new file mode 100644 index 0000000..a72e6c3 --- /dev/null +++ b/internal/ratelimit/redis.go @@ -0,0 +1,98 @@ +package ratelimit + +import ( + "context" + "fmt" + "time" +) + +// ScriptRunner is the slice of a Redis client this package needs. Narrow on +// purpose: it keeps go-redis out of this file's imports, which means the window +// logic below is testable without a Redis server. See redis_adapter.go for the +// ten-line bridge to *redis.Client. +type ScriptRunner interface { + Eval(ctx context.Context, script string, keys []string, args ...any) (any, error) +} + +// INCR and EXPIRE must be one atomic step. Done as two commands, a crash +// between them leaves a counter with no TTL, which silently becomes a permanent +// ban for that key. +// +// Returns {count, ttl_ms}. +const allowScript = ` +local current = redis.call('INCR', KEYS[1]) +if current == 1 then + redis.call('PEXPIRE', KEYS[1], ARGV[1]) +end +return {current, redis.call('PTTL', KEYS[1])} +` + +// Redis is a fixed-window limiter. +// +// Fixed window has a known flaw: a client can send `limit` requests at the end +// of one window and `limit` more at the start of the next, briefly achieving +// twice the rate. A sliding window or token bucket avoids it at the cost of +// more state. For protecting a backend from a runaway bot instance the simple +// version is enough, and the burst is bounded at 2x rather than unbounded. +type Redis struct { + Client ScriptRunner + Prefix string // namespace, so SCAN can find these keys when debugging +} + +func NewRedis(c ScriptRunner) *Redis { + return &Redis{Client: c, Prefix: "rl:"} +} + +func (r *Redis) Allow(ctx context.Context, key string, limit int, dur time.Duration) (Result, error) { + if limit <= 0 { + return Result{Allowed: true}, nil + } + + raw, err := r.Client.Eval(ctx, allowScript, + []string{r.Prefix + key}, dur.Milliseconds()) + if err != nil { + return Result{}, fmt.Errorf("ratelimit: eval: %w", err) + } + + vals, ok := raw.([]any) + if !ok || len(vals) != 2 { + return Result{}, fmt.Errorf("ratelimit: unexpected script result %T", raw) + } + + count, ok1 := toInt64(vals[0]) + ttlMS, ok2 := toInt64(vals[1]) + + if !ok1 || !ok2 { + return Result{}, fmt.Errorf("ratelimit: non-integer script result") + } + + remaining := limit - int(count) + if remaining < 0 { + remaining = 0 + } + + retry := time.Duration(ttlMS) * time.Millisecond + if ttlMS < 0 { + retry = dur + } + + return Result{ + Allowed: int(count) <= limit, + Limit: limit, + Remaining: remaining, + RetryAfter: retry, + }, nil +} + +// Redis integers arrive as int64 through go-redis, but a fake or a future +// client version may hand back int. Accept both rather than panicking. +func toInt64(v any) (int64, bool) { + switch n := v.(type) { + case int64: + return n, true + case int: + return int64(n), true + default: + return 0, false + } +} diff --git a/internal/ratelimit/redis_adapter.go b/internal/ratelimit/redis_adapter.go new file mode 100644 index 0000000..38399be --- /dev/null +++ b/internal/ratelimit/redis_adapter.go @@ -0,0 +1,20 @@ +package ratelimit + +import ( + "context" + + "github.com/redis/go-redis/v9" +) + +// GoRedis bridges *redis.Client to the ScriptRunner interface. +// +// Kept in its own file, and as thin as possible, because it is the one piece +// here that cannot be exercised without a live Redis: everything in redis.go is +// tested against a fake. +type GoRedis struct { + Client *redis.Client +} + +func (g GoRedis) Eval(ctx context.Context, script string, keys []string, args ...any) (any, error) { + return g.Client.Eval(ctx, script, keys, args...).Result() +}