317 lines
11 KiB
Python
317 lines
11 KiB
Python
"""
|
|
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
|
|
),
|
|
}
|