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:
+68
-40
@@ -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"))
|
||||
|
||||
Reference in New Issue
Block a user