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:
@@ -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
@@ -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"))
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
);
|
||||
Reference in New Issue
Block a user