diff --git a/internal/api/archiver.go b/internal/api/archiver.go index c014bae..25a0e5f 100644 --- a/internal/api/archiver.go +++ b/internal/api/archiver.go @@ -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. diff --git a/internal/api/auth.go b/internal/api/auth.go index 17f94fc..e161b51 100644 --- a/internal/api/auth.go +++ b/internal/api/auth.go @@ -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")) diff --git a/internal/api/device.go b/internal/api/device.go index 80e148c..d694b50 100644 --- a/internal/api/device.go +++ b/internal/api/device.go @@ -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 { diff --git a/internal/api/oidc.go b/internal/api/oidc.go index c38e317..fb121ac 100644 --- a/internal/api/oidc.go +++ b/internal/api/oidc.go @@ -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 { diff --git a/internal/api/rate_limiter_test.go b/internal/api/rate_limiter_test.go new file mode 100644 index 0000000..49831ae --- /dev/null +++ b/internal/api/rate_limiter_test.go @@ -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") + } +} diff --git a/internal/api/router.go b/internal/api/router.go index 3540795..cfe27b4 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -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) diff --git a/internal/api/signup.go b/internal/api/signup.go index 7bdbd2c..9b32e83 100644 --- a/internal/api/signup.go +++ b/internal/api/signup.go @@ -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 } diff --git a/internal/db/migrations/016_rate_limit_counters.sql b/internal/db/migrations/016_rate_limit_counters.sql new file mode 100644 index 0000000..5553401 --- /dev/null +++ b/internal/db/migrations/016_rate_limit_counters.sql @@ -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 +);