""" 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). 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) """ import logging import os 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 _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) """ _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.""" _cleanup_expired() _ip_attempts[ip] = [(time.time(), True)] def is_rate_limited(ip: str) -> bool: """Check if an IP has exceeded the rate limit.""" _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) """ _cleanup_expired() key = account.lower() _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.""" _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.""" _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).""" _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 ), }