"""JWT authentication module for Vacuum Wall. Handles token creation, validation, refresh, and blacklisting. Each user has their own JWT signing secret stored in the database. Configuration comes from config/auth/config.json. """ from __future__ import annotations import base64 import json import logging import secrets import time import uuid from pathlib import Path from typing import Any import jwt from lib.common import load_json from lib.db import ( Q_DELETE_EXPIRED_BLACKLIST, Q_DELETE_REFRESH_TOKEN, Q_INSERT_BLACKLIST, Q_SELECT_BLACKLIST, Q_SELECT_REFRESH_TOKEN, Q_SELECT_USER_JWT_SECRET, Q_UPDATE_JWT_SECRET, Q_UPSERT_REFRESH_TOKEN, get_db, ) logger = logging.getLogger(__name__) AUTH_CONFIG_PATH = ( Path(__file__).resolve().parent.parent / "config" / "auth" / "config.json" ) _DEFAULT_JWT_CONFIG = { "access_token_ttl": 900, "refresh_token_ttl": 604800, "algorithm": "HS256", } def _get_jwt_config() -> dict[str, Any]: """Load JWT configuration from auth config.""" raw = load_json(AUTH_CONFIG_PATH) return raw.get("jwt", _DEFAULT_JWT_CONFIG) def get_user_jwt_secret(username: str) -> str | None: """Return the JWT signing secret for *username*, or ``None`` if not found.""" db = get_db() rows = db.query(Q_SELECT_USER_JWT_SECRET, (username,)) if not rows: return None return rows[0]["jwt_secret"] def rotate_user_secret(username: str) -> None: """Rotate the JWT secret for *username*, invalidating all their existing tokens. Used when a user's password is changed to ensure all prior sessions are immediately terminated regardless of token expiration. """ new_secret = secrets.token_urlsafe(32) db = get_db() db.run(Q_UPDATE_JWT_SECRET, (new_secret, username)) logger.warning( "JWT secret rotated for %r — their existing tokens are now invalid", username ) def get_access_ttl() -> int: """Return access token TTL in seconds.""" return _get_jwt_config().get("access_token_ttl", 900) def get_refresh_ttl() -> int: """Return refresh token TTL in seconds.""" return _get_jwt_config().get("refresh_token_ttl", 604800) def get_algorithm() -> str: """Return the JWT algorithm.""" return _get_jwt_config().get("algorithm", "HS256") def generate_access_token( username: str, permissions: dict[str, str], session_id: str | None = None, ) -> str: """Generate a new access token for *username*. Args: username: The authenticated username. permissions: Dict mapping subsystem names to permission levels. session_id: Optional session binding ID. Included in the token payload so the Flask middleware can tie the token to the browser session that created it. Returns: JWT token string. Raises: RuntimeError: If JWT secret is not configured. """ secret = get_user_jwt_secret(username) if not secret: raise RuntimeError(f"JWT secret not configured for user {username!r}") algorithm = get_algorithm() now = int(time.time()) payload = { "sub": username, "exp": now + get_access_ttl(), "iat": now, "jti": str(uuid.uuid4()), "type": "access", "permissions": permissions, "session_id": session_id or secrets.token_urlsafe(16), } return jwt.encode(payload, secret, algorithm=algorithm) def generate_refresh_token(username: str) -> str: """Generate a new refresh token for *username*. Args: username: The authenticated username. Returns: JWT refresh token string. Raises: RuntimeError: If JWT secret is not configured. """ secret = get_user_jwt_secret(username) if not secret: raise RuntimeError(f"JWT secret not configured for user {username!r}") algorithm = get_algorithm() now = int(time.time()) payload = { "sub": username, "exp": now + get_refresh_ttl(), "iat": now, "jti": str(uuid.uuid4()), "type": "refresh", } return jwt.encode(payload, secret, algorithm=algorithm) def generate_tokens(username: str, permissions: dict[str, str]) -> dict[str, str]: """Generate both access and refresh tokens. Args: username: The authenticated username. permissions: Dict mapping subsystem names to permission levels. Returns: Dict with ``access_token`` and ``refresh_token`` keys. """ session_id = secrets.token_urlsafe(16) access_token = generate_access_token(username, permissions, session_id) refresh_token = generate_refresh_token(username) _persist_refresh_token(username, refresh_token) return { "access_token": access_token, "refresh_token": refresh_token, "session_id": session_id, } def _persist_refresh_token(username: str, refresh_token: str) -> None: """Persist the active refresh token JTI for *username* in the database. One row per user — replaces any existing entry on upsert. Args: username: The username. refresh_token: The JWT refresh token string. """ payload = decode_token(refresh_token) if payload is None: return jti = payload.get("jti") if not jti: return issued_at = payload.get("iat", int(time.time())) db = get_db() db.run(Q_UPSERT_REFRESH_TOKEN, (username, jti, issued_at)) def blacklist_active_refresh_token(username: str) -> None: """Blacklist the user's currently active refresh token from the database. Looks up the stored JTI for *username*, blacklists it, and removes the database entry. Safe to call when no token is registered — the query will simply return no rows. Args: username: The username. """ db = get_db() rows = db.query(Q_SELECT_REFRESH_TOKEN, (username,)) if rows: blacklist_token(rows[0]["jti"], token_type="refresh") db.run(Q_DELETE_REFRESH_TOKEN, (username,)) def clear_active_refresh_token(username: str) -> None: """Remove the user's stored refresh token entry without blacklisting. Used when the refresh token has already been blacklisted (e.g., during a successful refresh rotation). Args: username: The username. """ db = get_db() db.run(Q_DELETE_REFRESH_TOKEN, (username,)) def _extract_unverified_sub(token_string: str) -> str | None: """Extract the ``sub`` claim from a JWT payload without signature verification. The JWT payload is the second segment (dot-separated), base64url-encoded JSON. This is safe because we are NOT trusting the claim value — we use it solely to look up the user's secret for proper verification. Args: token_string: The JWT token string. Returns: The ``sub`` claim value, or ``None`` if the token is malformed. """ try: parts = token_string.split(".") if len(parts) != 3: return None payload_b64 = parts[1] # Add padding padding = 4 - len(payload_b64) % 4 if padding != 4: payload_b64 += "=" * padding payload_json = base64.urlsafe_b64decode(payload_b64) payload = json.loads(payload_json) return payload.get("sub") except (ValueError, json.JSONDecodeError, UnicodeDecodeError): return None def decode_token(token_string: str) -> dict[str, Any] | None: """Decode and validate a JWT token using the user's secret. Extracts the ``sub`` claim from the unverified payload to look up the correct per-user signing secret, then verifies the signature. Args: token_string: The JWT token string (without Bearer prefix). Returns: Payload dict if valid, None if invalid/expired or user not found. """ sub = _extract_unverified_sub(token_string) if not sub: return None secret = get_user_jwt_secret(sub) if not secret: return None algorithm = get_algorithm() try: payload = jwt.decode(token_string, secret, algorithms=[algorithm]) return payload except (jwt.ExpiredSignatureError, jwt.InvalidTokenError): return None def validate_token( token_string: str, token_type: str = "access", session_id: str | None = None, ) -> dict[str, Any] | None: """Validate a JWT token and check it against the blacklist. Args: token_string: The JWT token string. token_type: Expected token type ("access" or "refresh"). session_id: If provided, must match the ``session_id`` claim in the token payload. Acts as session binding — prevents a stolen token from being used by an attacker who doesn't also possess the matching session ID. Returns: Payload dict including permissions, or None if invalid/blacklisted. """ payload = decode_token(token_string) if payload is None: return None if payload.get("type") != token_type: return None if session_id and payload.get("session_id") != session_id: return None jti = payload.get("jti") if jti and is_blacklisted(jti): return None return payload def blacklist_token(jti: str, token_type: str = "access") -> None: """Add a JTI to the blacklist. Args: jti: Token UUID to blacklist. token_type: Token type ("access" or "refresh"). Defaults to "access". The blacklist expiry is set to current time + the token type's TTL. """ db = get_db() ttl = get_refresh_ttl() if token_type == "refresh" else get_access_ttl() db.run(Q_INSERT_BLACKLIST, (jti, token_type, int(time.time()) + ttl)) def is_blacklisted(jti: str) -> bool: """Check if a JTI is blacklisted. Args: jti: Token UUID to check. Returns: True if the token has been blacklisted. """ db = get_db() rows = db.query(Q_SELECT_BLACKLIST, (jti,)) return len(rows) > 0 def blacklist_expired() -> None: """Remove expired entries from the blacklist.""" db = get_db() now = int(time.time()) db.run(Q_DELETE_EXPIRED_BLACKLIST, (now,)) class RateLimiter: """Simple sliding-window rate limiter for login attempts. Maintains per-key attempt timestamps and rejects requests that exceed the allowed count within the window. """ def __init__(self, max_attempts: int = 5, window_seconds: int = 300) -> None: self.max_attempts = max_attempts self.window = window_seconds self.attempts: dict[str, list[float]] = {} def is_allowed(self, key: str) -> bool: """Check if a request from *key* is allowed. Args: key: Identifier for the rate limit bucket (e.g., username or IP). Returns: True if the request is allowed, False if rate limited. """ now = time.time() cutoff = now - self.window timestamps = self.attempts.get(key, []) # Clean old entries self.attempts[key] = [t for t in timestamps if t > cutoff] if len(self.attempts[key]) >= self.max_attempts: return False self.attempts[key].append(now) return True def cleanup(self) -> None: """Remove expired entries from all buckets.""" now = time.time() cutoff = now - self.window for key in list(self.attempts): self.attempts[key] = [t for t in self.attempts[key] if t > cutoff] if not self.attempts[key]: del self.attempts[key] # Global rate limiters _login_limiter = RateLimiter(max_attempts=10, window_seconds=300) _webauthn_limiter = RateLimiter(max_attempts=5, window_seconds=600) def check_login_rate(username: str) -> bool: """Check if login is rate-limited for the given username. Args: username: The login attempt username. Returns: True if the attempt is allowed, False if rate limited. """ return _login_limiter.is_allowed(username) def check_webauthn_rate(username: str) -> bool: """Check if WebAuthn authentication is rate-limited for the given username. Args: username: The WebAuthn attempt username. Returns: True if the attempt is allowed, False if rate limited. """ return _webauthn_limiter.is_allowed(username)