""" In-memory rate limiter for authentication endpoints. Tracks failed attempts per IP **and** per account with automatic cleanup of expired entries. The per-IP budget stops a single source; the per-account budget (BUG-031) still throttles an attacker who rotates IPs. It complements the per-account lockout in ``user_store.py``. .. note:: The counters live in process memory only. They are **not** shared between multiple workers/containers and are lost on restart. For a multi-node deployment, front this service with a shared store (Redis) or a single worker. This limitation is intentional and documented (BUG-031). Opt-in persistence (ROADMAP #85 T10b) : if ``OBSIGATE_RATELIMIT_DB`` points to a SQLite file, counters are stored there instead (WAL mode, one short connection per call — safe across threads, processes and restarts sharing the same file). Semantics (windows, budgets, success reset) are identical to the in-memory store, which remains the default when the variable is unset. Configuration via environment variables: OBSIGATE_LOGIN_MAX_ATTEMPTS Max failures per IP (default: 10) OBSIGATE_ACCOUNT_MAX_ATTEMPTS Max failures per account (default: 10) OBSIGATE_LOGIN_WINDOW_SECONDS Lockout window in seconds (default: 900) OBSIGATE_RATELIMIT_DB SQLite file for shared/persistent counters (default: unset = memory) """ import logging import os import sqlite3 import threading import time from collections import defaultdict logger = logging.getLogger("obsigate.ratelimit") # --- Configuration --- MAX_ATTEMPTS = int(os.environ.get("OBSIGATE_LOGIN_MAX_ATTEMPTS", "10")) ACCOUNT_MAX_ATTEMPTS = int(os.environ.get("OBSIGATE_ACCOUNT_MAX_ATTEMPTS", "10")) WINDOW_SECONDS = int(os.environ.get("OBSIGATE_LOGIN_WINDOW_SECONDS", "900")) # 15 min # --- In-memory stores: {key: [(timestamp, success_bool), ...]} --- _ip_attempts: dict[str, list] = defaultdict(list) _account_attempts: dict[str, list] = defaultdict(list) _last_cleanup = time.time() CLEANUP_INTERVAL = 60 # seconds def _db_path() -> str | None: """SQLite file for shared counters, or ``None`` for the in-memory store.""" path = os.environ.get("OBSIGATE_RATELIMIT_DB", "").strip() return path or None def _db_connect(path: str) -> sqlite3.Connection: """Open a short-lived connection (WAL + busy timeout for concurrent workers).""" _db_ensure_schema(path) conn = sqlite3.connect(path, timeout=10.0) conn.execute("PRAGMA busy_timeout=10000") return conn _schema_ready: set[str] = set() _schema_lock = threading.Lock() def _db_ensure_schema(path: str) -> None: """Create the store schema once per file (DDL under a process-wide lock).""" with _schema_lock: if path in _schema_ready: return conn = sqlite3.connect(path, timeout=10.0) try: conn.execute("PRAGMA journal_mode=WAL") conn.execute( "CREATE TABLE IF NOT EXISTS attempts" " (kind TEXT NOT NULL, key TEXT NOT NULL, ts REAL NOT NULL, success INTEGER NOT NULL)" ) conn.execute( "CREATE INDEX IF NOT EXISTS idx_attempts_kind_key_ts" " ON attempts (kind, key, ts)" ) conn.commit() finally: conn.close() _schema_ready.add(path) def _db_write(fn, *args): """Run a write op, retrying once on lock contention (concurrent workers).""" try: return fn(*args) except sqlite3.OperationalError as e: if "locked" not in str(e).lower(): raise time.sleep(0.05) return fn(*args) def _db_prune(conn: sqlite3.Connection, cutoff: float) -> None: """Drop expired entries (best-effort cap on disk growth).""" conn.execute("DELETE FROM attempts WHERE ts <= ?", (cutoff,)) def _db_record(kind: str, key: str, success: bool) -> int: """Record one attempt in SQLite; return the live failure count.""" path = _db_path() assert path is not None now = time.time() cutoff = now - WINDOW_SECONDS def _write() -> int: with _db_connect(path) as conn: _db_prune(conn, cutoff) if success: # Mirror the in-memory reset: replace history with one success. conn.execute("DELETE FROM attempts WHERE kind = ? AND key = ?", (kind, key)) conn.execute( "INSERT INTO attempts (kind, key, ts, success) VALUES (?, ?, ?, ?)", (kind, key, now, int(success)), ) conn.commit() (failures,) = conn.execute( "SELECT COUNT(*) FROM attempts WHERE kind = ? AND key = ? AND ts > ? AND success = 0", (kind, key, cutoff), ).fetchone() return failures return _db_write(_write) def _db_failures(kind: str, key: str) -> int: """Live failure count in SQLite (expired entries never count).""" path = _db_path() assert path is not None cutoff = time.time() - WINDOW_SECONDS with _db_connect(path) as conn: (failures,) = conn.execute( "SELECT COUNT(*) FROM attempts WHERE kind = ? AND key = ? AND ts > ? AND success = 0", (kind, key, cutoff), ).fetchone() return failures def _db_tracked(kind: str) -> int: """Number of distinct keys ever seen for one budget (SQLite).""" path = _db_path() assert path is not None with _db_connect(path) as conn: (n,) = conn.execute( "SELECT COUNT(DISTINCT key) FROM attempts WHERE kind = ?", (kind,) ).fetchone() return n def _db_limited_count(kind: str, max_attempts: int) -> int: """Number of keys currently over budget (SQLite).""" path = _db_path() assert path is not None cutoff = time.time() - WINDOW_SECONDS with _db_connect(path) as conn: rows = conn.execute( "SELECT key, COUNT(*) FROM attempts" " WHERE kind = ? AND ts > ? AND success = 0 GROUP BY key", (kind, cutoff), ).fetchall() return sum(1 for _, n in rows if n >= max_attempts) def _prune(store: dict[str, list], cutoff: float) -> None: """Drop expired entries from one store in place.""" expired = [] for key, attempts in store.items(): store[key] = [a for a in attempts if a[0] > cutoff] if not store[key]: expired.append(key) for key in expired: del store[key] def _cleanup_expired(): """Remove entries older than the window from both stores.""" global _last_cleanup now = time.time() if now - _last_cleanup < CLEANUP_INTERVAL: return _last_cleanup = now cutoff = now - WINDOW_SECONDS _prune(_ip_attempts, cutoff) _prune(_account_attempts, cutoff) def record_failure(ip: str) -> tuple[int, int]: """Record a failed login attempt from an IP. Returns: (current_failure_count, remaining_attempts) """ if _db_path() is not None: failures = _db_record("ip", ip, False) remaining = max(0, MAX_ATTEMPTS - failures) if failures >= MAX_ATTEMPTS: logger.warning(f"IP {ip} rate-limited after {failures} failed logins") return failures, remaining _cleanup_expired() _ip_attempts[ip].append((time.time(), False)) failures = sum(1 for _, success in _ip_attempts[ip] if not success) remaining = max(0, MAX_ATTEMPTS - failures) if failures >= MAX_ATTEMPTS: logger.warning(f"IP {ip} rate-limited after {failures} failed logins") return failures, remaining def record_success(ip: str): """Clear rate limit state for an IP after successful login.""" if _db_path() is not None: _db_record("ip", ip, True) return _cleanup_expired() _ip_attempts[ip] = [(time.time(), True)] def is_rate_limited(ip: str) -> bool: """Check if an IP has exceeded the rate limit.""" if _db_path() is not None: return _db_failures("ip", ip) >= MAX_ATTEMPTS _cleanup_expired() failures = sum(1 for _, success in _ip_attempts.get(ip, []) if not success) return failures >= MAX_ATTEMPTS def record_account_failure(account: str) -> tuple[int, int]: """Record a failed attempt for an account, regardless of source IP. Returns: (current_failure_count, remaining_attempts) """ key = account.lower() if _db_path() is not None: failures = _db_record("account", key, False) remaining = max(0, ACCOUNT_MAX_ATTEMPTS - failures) if failures >= ACCOUNT_MAX_ATTEMPTS: logger.warning(f"Account {account} rate-limited after {failures} failed attempts") return failures, remaining _cleanup_expired() _account_attempts[key].append((time.time(), False)) failures = sum(1 for _, success in _account_attempts[key] if not success) remaining = max(0, ACCOUNT_MAX_ATTEMPTS - failures) if failures >= ACCOUNT_MAX_ATTEMPTS: logger.warning(f"Account {account} rate-limited after {failures} failed attempts") return failures, remaining def record_account_success(account: str): """Clear the per-account rate limit state after a successful login.""" if _db_path() is not None: _db_record("account", account.lower(), True) return _cleanup_expired() _account_attempts[account.lower()] = [(time.time(), True)] def is_account_rate_limited(account: str) -> bool: """Check if an account has exceeded the per-account rate limit.""" if _db_path() is not None: return _db_failures("account", account.lower()) >= ACCOUNT_MAX_ATTEMPTS _cleanup_expired() failures = sum( 1 for _, success in _account_attempts.get(account.lower(), []) if not success ) return failures >= ACCOUNT_MAX_ATTEMPTS def get_status(ip: str | None = None) -> dict: """Get rate limit status for an IP (for diagnostics).""" if _db_path() is not None: if ip: failures = _db_failures("ip", ip) return { "ip": ip, "failures": failures, "max": MAX_ATTEMPTS, "limited": failures >= MAX_ATTEMPTS, "window_seconds": WINDOW_SECONDS, } return { "tracked_ips": _db_tracked("ip"), "tracked_accounts": _db_tracked("account"), "max_attempts": MAX_ATTEMPTS, "account_max_attempts": ACCOUNT_MAX_ATTEMPTS, "window_seconds": WINDOW_SECONDS, "limited_ips": _db_limited_count("ip", MAX_ATTEMPTS), } _cleanup_expired() if ip: attempts = _ip_attempts.get(ip, []) failures = sum(1 for _, s in attempts if not s) return { "ip": ip, "failures": failures, "max": MAX_ATTEMPTS, "limited": failures >= MAX_ATTEMPTS, "window_seconds": WINDOW_SECONDS, } return { "tracked_ips": len(_ip_attempts), "tracked_accounts": len(_account_attempts), "max_attempts": MAX_ATTEMPTS, "account_max_attempts": ACCOUNT_MAX_ATTEMPTS, "window_seconds": WINDOW_SECONDS, "limited_ips": sum( 1 for ip_addr in _ip_attempts if sum(1 for _, s in _ip_attempts[ip_addr] if not s) >= MAX_ATTEMPTS ), }