Back the login/signup/OIDC/device rate limiters with Postgres

Part of the same security-hardening pass as the body-size/header commit.
loginLimiter was an in-memory, per-process sync.Mutex+map -- fine for one
replica, but charts/terdut-server/values.yaml has set replicaCount: 2 in
production since v0.37.0. Each pod counted only its own traffic, so every
limit it guarded (failed logins, sign-ups, OIDC/device-login starts) was
effectively twice as generous as the constants say, not just in theory.

loginLimiter now stores its counters in a new rate_limit_counters table
(migration 016) instead of a map; blocked/fail/clear take a context and
query/upsert/delete a row keyed by the same strings callers already used
(username, client address, "signup:"+address, ...). Semantics are
unchanged -- a fixed window that resets rather than slides -- so no call
site's behavior changes, only where the count lives. Added purgeRateLimits
to the sweeper, alongside purgeSessions/purgeAckTokens, so expired windows
don't accumulate.

New internal (package api) tests in rate_limiter_test.go cover the basic
behavior plus the regression this exists to fix: two loginLimiter values
sharing one database, standing in for two replicas, now see one combined
count instead of each keeping their own.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
Niklas Ye
2026-10-07 22:03:51 +02:00
parent 7cd6fbf571
commit b82c10acf4
8 changed files with 255 additions and 49 deletions
+1
View File
@@ -76,6 +76,7 @@ func Sweep(ctx context.Context, db *sql.DB, archiveAfter, staleAfter time.Durati
archiveResolvedIncidents(ctx, db, archiveAfter)
purgeAckTokens(ctx, db)
purgeSessions(ctx, db)
purgeRateLimits(ctx, db)
}
// expireStale resolves firing alerts that Alertmanager has stopped refreshing.
+68 -40
View File
@@ -45,57 +45,85 @@ var dummyHash = sync.OnceValue(func() []byte {
return h
})
// loginLimiter counts failed logins in a fixed window, per username and per
// client address. The username limit is what stops guessing one account; the
// address limit is looser because every user behind the same gateway or NAT
// shares it.
// loginLimiter counts failed logins (and other unauthenticated attempts:
// sign-up, OIDC/device start) in a fixed window, per key — a username, a
// client address, or both, depending on the caller.
//
// Backed by Postgres rather than an in-memory map: this server runs more
// than one replica in production (v0.37.0), and a counter that only ever
// sees its own pod's traffic would quietly let every limit through
// multiplied by the replica count — two loginLimiter values pointed at the
// same db, standing in for two replicas, now share exactly one count per
// key instead of each keeping their own.
//
// The window resets rather than slides, the same behavior the in-memory
// version it replaces had: once a key's window is older than loginWindow,
// the next fail() starts a fresh one instead of extending the stale one.
type loginLimiter struct {
mu sync.Mutex
failures map[string]*loginWindowCount
db *sql.DB
}
type loginWindowCount struct {
start time.Time
n int
func newLoginLimiter(db *sql.DB) *loginLimiter {
return &loginLimiter{db: db}
}
func newLoginLimiter() *loginLimiter {
return &loginLimiter{failures: map[string]*loginWindowCount{}}
}
func (l *loginLimiter) blocked(key string, max int) bool {
l.mu.Lock()
defer l.mu.Unlock()
c, ok := l.failures[key]
if !ok || time.Since(c.start) > loginWindow {
func (l *loginLimiter) blocked(ctx context.Context, key string, max int) bool {
cutoff := time.Now().Unix() - int64(loginWindow.Seconds())
var count int
err := l.db.QueryRowContext(ctx, `
SELECT count FROM rate_limit_counters
WHERE key = $1 AND window_start > $2`,
key, cutoff,
).Scan(&count)
if err != nil {
// No row (never failed, or its window already expired): not blocked.
// A real query error fails the same way — a rate limiter that locks
// everyone out during a brief database hiccup is worse than one that
// is briefly too generous.
return false
}
return c.n >= max
return count >= max
}
func (l *loginLimiter) fail(keys ...string) {
l.mu.Lock()
defer l.mu.Unlock()
now := time.Now()
for k, c := range l.failures {
if now.Sub(c.start) > loginWindow {
delete(l.failures, k)
}
}
func (l *loginLimiter) fail(ctx context.Context, keys ...string) {
now := time.Now().Unix()
windowSecs := int64(loginWindow.Seconds())
for _, key := range keys {
c, ok := l.failures[key]
if !ok {
c = &loginWindowCount{start: now}
l.failures[key] = c
if _, err := l.db.ExecContext(ctx, `
INSERT INTO rate_limit_counters (key, window_start, count)
VALUES ($1, $2, 1)
ON CONFLICT (key) DO UPDATE SET
window_start = CASE WHEN rate_limit_counters.window_start <= $2 - $3
THEN $2 ELSE rate_limit_counters.window_start END,
count = CASE WHEN rate_limit_counters.window_start <= $2 - $3
THEN 1 ELSE rate_limit_counters.count + 1 END`,
key, now, windowSecs,
); err != nil {
log.Printf("rate limiter: record failure for %q: %v", key, err)
}
c.n++
}
}
func (l *loginLimiter) clear(key string) {
l.mu.Lock()
defer l.mu.Unlock()
delete(l.failures, key)
func (l *loginLimiter) clear(ctx context.Context, key string) {
if _, err := l.db.ExecContext(ctx, "DELETE FROM rate_limit_counters WHERE key = $1", key); err != nil {
log.Printf("rate limiter: clear %q: %v", key, err)
}
}
// purgeRateLimits deletes rate-limit windows that have expired, from the
// sweeper — otherwise every distinct username and address this server has
// ever seen a failed attempt from would stay a row forever.
func purgeRateLimits(ctx context.Context, db *sql.DB) {
cutoff := time.Now().Unix() - int64(loginWindow.Seconds())
res, err := db.ExecContext(ctx,
"DELETE FROM rate_limit_counters WHERE window_start <= $1", cutoff)
if err != nil {
log.Printf("sweeper: purge rate limit counters: %v", err)
return
}
if n, _ := res.RowsAffected(); n > 0 {
log.Printf("sweeper: purged %d expired rate limit counter(s)", n)
}
}
// clientAddr is the address a login is counted against. Behind the gateway
@@ -198,7 +226,7 @@ func handleLogin(db *sql.DB, limiter *loginLimiter, publicURL string) http.Handl
userKey := "user:" + strings.ToLower(username)
addrKey := "addr:" + clientAddr(r)
if limiter.blocked(userKey, loginMaxPerUser) || limiter.blocked(addrKey, loginMaxPerAddr) {
if limiter.blocked(r.Context(), userKey, loginMaxPerUser) || limiter.blocked(r.Context(), addrKey, loginMaxPerAddr) {
w.Header().Set("Retry-After", strconv.Itoa(int(loginWindow.Seconds())))
respond(w, http.StatusTooManyRequests, errResp("too many failed attempts, try again later"))
return
@@ -220,11 +248,11 @@ func handleLogin(db *sql.DB, limiter *loginLimiter, publicURL string) http.Handl
}
match := bcrypt.CompareHashAndPassword(stored, []byte(req.Password)) == nil
if !match || !hash.Valid {
limiter.fail(userKey, addrKey)
limiter.fail(r.Context(), userKey, addrKey)
respond(w, http.StatusUnauthorized, errResp("invalid username or password"))
return
}
limiter.clear(userKey)
limiter.clear(r.Context(), userKey)
if err := startSession(w, r, db, userID, publicURL); err != nil {
respond(w, http.StatusInternalServerError, errResp("internal error"))
+2 -2
View File
@@ -72,12 +72,12 @@ func normalizeUserCode(s string) string {
func handleDeviceStart(db *sql.DB, limiter *loginLimiter, publicURL string) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
addrKey := "device:" + clientAddr(r)
if limiter.blocked(addrKey, deviceStartMaxPerAddr) {
if limiter.blocked(r.Context(), addrKey, deviceStartMaxPerAddr) {
w.Header().Set("Retry-After", strconv.Itoa(int(loginWindow.Seconds())))
respond(w, http.StatusTooManyRequests, errResp("too many sign-in attempts, try again later"))
return
}
limiter.fail(addrKey)
limiter.fail(r.Context(), addrKey)
deviceCode, deviceHash, err := randomToken()
if err != nil {
+2 -2
View File
@@ -121,12 +121,12 @@ func ssoRedirect(w http.ResponseWriter, r *http.Request, code ssoError) {
func handleOIDCLogin(db *sql.DB, prov *oidc.Provider, limiter *loginLimiter, publicURL string) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
addrKey := "oidc:" + clientAddr(r)
if limiter.blocked(addrKey, oidcStartMaxPerAddr) {
if limiter.blocked(r.Context(), addrKey, oidcStartMaxPerAddr) {
w.Header().Set("Retry-After", strconv.Itoa(int(loginWindow.Seconds())))
respond(w, http.StatusTooManyRequests, errResp("too many sign-in attempts, try again later"))
return
}
limiter.fail(addrKey)
limiter.fail(r.Context(), addrKey)
state, stateHash, err := randomToken()
if err != nil {
+161
View File
@@ -0,0 +1,161 @@
package api
// This file is internal (package api, not api_test) because loginLimiter and
// its blocked/fail/clear methods are unexported, and TestLoginLimiter_SharedAcrossReplicas
// specifically needs to construct two separate loginLimiter values pointed at
// one database — standing in for two replicas — which only this package can
// do. It duplicates testdb_test.go's newTestDB/withSearchPath rather than
// importing them: those live in the separate api_test package, compiled from
// this directory's external test files, and are not visible here. Same
// reasoning as advisory_lock_test.go, which makes the same trade for the
// same reason.
import (
"context"
"database/sql"
"fmt"
"net/url"
"os"
"strings"
"testing"
"git.ryuvia.com/niklas/terdut-server/internal/db"
_ "github.com/jackc/pgx/v5/stdlib"
)
var rateLimiterSchemaSeq int
// rateLimiterTestDB returns a migrated database private to this test.
func rateLimiterTestDB(t *testing.T) *sql.DB {
t.Helper()
dsn := os.Getenv("TERDUT_TEST_DSN")
if dsn == "" {
t.Fatalf("TERDUT_TEST_DSN is not set: these tests need Postgres.\n" +
"Run `make test-db` for a local one, then\n" +
" export TERDUT_TEST_DSN=postgres://terdut:terdut@localhost:5432/terdut_test?sslmode=disable")
}
rateLimiterSchemaSeq++
schema := fmt.Sprintf("test_rl_%d_%d", os.Getpid(), rateLimiterSchemaSeq)
admin, err := sql.Open("pgx", dsn)
if err != nil {
t.Fatalf("connect to TERDUT_TEST_DSN: %v", err)
}
defer admin.Close()
if _, err := admin.Exec("CREATE SCHEMA " + schema); err != nil {
t.Fatalf("create schema %s: %v", schema, err)
}
database, err := db.Open(rateLimiterWithSearchPath(dsn, schema))
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.Migrate(database); err != nil {
t.Fatalf("migrate: %v", err)
}
t.Cleanup(func() {
database.Close()
cleanup, err := sql.Open("pgx", dsn)
if err != nil {
return
}
defer cleanup.Close()
if _, err := cleanup.Exec("DROP SCHEMA " + schema + " CASCADE"); err != nil {
t.Logf("drop schema %s: %v", schema, err)
}
})
return database
}
func rateLimiterWithSearchPath(dsn, schema string) string {
opt := "-csearch_path=" + schema
if strings.HasPrefix(dsn, "postgres://") || strings.HasPrefix(dsn, "postgresql://") {
u, err := url.Parse(dsn)
if err == nil {
q := u.Query()
q.Set("options", opt)
u.RawQuery = q.Encode()
return u.String()
}
}
return dsn + " options='" + opt + "'"
}
func TestLoginLimiter_BlocksAtMax(t *testing.T) {
database := rateLimiterTestDB(t)
ctx := context.Background()
l := newLoginLimiter(database)
for range 3 {
if l.blocked(ctx, "k", 3) {
t.Fatal("blocked before reaching max")
}
l.fail(ctx, "k")
}
if !l.blocked(ctx, "k", 3) {
t.Fatal("not blocked after reaching max")
}
}
func TestLoginLimiter_ClearResetsTheCount(t *testing.T) {
database := rateLimiterTestDB(t)
ctx := context.Background()
l := newLoginLimiter(database)
l.fail(ctx, "k")
l.fail(ctx, "k")
l.clear(ctx, "k")
if l.blocked(ctx, "k", 1) {
t.Fatal("still blocked after clear")
}
}
func TestLoginLimiter_KeysAreIndependent(t *testing.T) {
database := rateLimiterTestDB(t)
ctx := context.Background()
l := newLoginLimiter(database)
l.fail(ctx, "a")
if l.blocked(ctx, "b", 1) {
t.Fatal("failing one key blocked an unrelated one")
}
}
// TestLoginLimiter_SharedAcrossReplicas is the regression test for the gap
// this migration closes: an in-memory limiter would let each replica count
// independently, so a caller hitting two different pods could rack up
// max*replicaCount failures before either one blocked. Two loginLimiter
// values sharing one database, standing in for two replicas behind the same
// load balancer, must instead see one combined count.
func TestLoginLimiter_SharedAcrossReplicas(t *testing.T) {
database := rateLimiterTestDB(t)
ctx := context.Background()
replicaA := newLoginLimiter(database)
replicaB := newLoginLimiter(database)
const max = 4
// Alternate which "replica" records the failure, as a real deployment
// would split requests across pods.
for i := range max {
replica := replicaA
if i%2 == 1 {
replica = replicaB
}
if replicaA.blocked(ctx, "k", max) || replicaB.blocked(ctx, "k", max) {
t.Fatalf("blocked after only %d of %d failures", i, max)
}
replica.fail(ctx, "k")
}
if !replicaA.blocked(ctx, "k", max) {
t.Fatal("replica A does not see the combined count as blocked")
}
if !replicaB.blocked(ctx, "k", max) {
t.Fatal("replica B does not see the combined count as blocked")
}
}
+3 -3
View File
@@ -23,9 +23,9 @@ func NewRouter(db *sql.DB, notify NotifyConfig, cfg config.Config, version strin
// One limiter each, both process-wide for the life of the router: login
// counts failed passwords, sign-up counts account creation, and mixing the
// two would let a burst of sign-ups lock somebody out of logging in.
loginLimit := newLoginLimiter()
signupLimiter := newLoginLimiter()
oidcLimit := newLoginLimiter()
loginLimit := newLoginLimiter(db)
signupLimiter := newLoginLimiter(db)
oidcLimit := newLoginLimiter(db)
r := chi.NewRouter()
r.Use(middleware.Logger)
+2 -2
View File
@@ -112,7 +112,7 @@ var errInviteUnusable = errors.New("invite is not usable")
func handleSignup(db *sql.DB, limiter *loginLimiter, publicURL string) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
addr := clientAddr(r)
if limiter.blocked("signup:"+addr, maxSignupsPerAddr) {
if limiter.blocked(r.Context(), "signup:"+addr, maxSignupsPerAddr) {
respond(w, http.StatusTooManyRequests, errResp("too many sign-ups from this address"))
return
}
@@ -148,7 +148,7 @@ func handleSignup(db *sql.DB, limiter *loginLimiter, publicURL string) http.Hand
var err error
inv, err = loadInvite(r.Context(), db, req.Invite)
if err != nil {
limiter.fail("signup:" + addr)
limiter.fail(r.Context(), "signup:"+addr)
respond(w, http.StatusForbidden, errResp("this invite link is not usable"))
return
}
@@ -0,0 +1,16 @@
-- Backs the rate limiters (failed logins, sign-ups, OIDC/device start) with
-- Postgres instead of an in-memory map, now that the server runs more than
-- one replica in production (v0.37.0): a counter that only ever sees its own
-- pod's traffic quietly let every one of these limits through multiplied by
-- the replica count.
--
-- window_start is the start of the current fixed window for key, in the same
-- "unix seconds" shape every other timestamp in this schema uses. The window
-- resets rather than slides, matching the in-memory limiter it replaces:
-- once a key's window is older than the limiter's window length, the next
-- failure starts a fresh one instead of extending the stale one.
CREATE TABLE rate_limit_counters (
key TEXT PRIMARY KEY,
window_start BIGINT NOT NULL,
count INT NOT NULL
);