Files
vacuum-wall/lib/auth.py
T
mteehan c7593f8a1e fix: close refresh token session binding bypass
Add require_session parameter to validate_token to enforce session_id
matching for refresh operations. Attacker with stolen refresh token can
no longer bypass session binding by omitting session_id from request.

Also adds backend guard against deleting builtin admin user (was only
blocked at Flask blueprint layer), and removes unused _ALL_RW variable.
2026-08-12 20:06:05 +00:00

479 lines
15 KiB
Python

"""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 random
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, session_id: str | None = None) -> str:
"""Generate a new refresh token for *username*.
Args:
username: The authenticated username.
session_id: Session binding ID included in the token payload.
When present, the refresh endpoint requires a matching session_id,
preventing a stolen refresh token from being usable without the
originating browser session.
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",
}
if session_id:
payload["session_id"] = session_id
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, session_id)
_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,
require_session: bool = False,
) -> 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: Must match the ``session_id`` claim in the token payload.
When provided, enforces session binding to prevent a stolen token
from being usable without the originating session. When ``None``,
the check is skipped unless ``require_session`` is True.
require_session: If True and the token payload contains a ``session_id``,
the request must provide a matching ``session_id``. Used by the
refresh handler to prevent session binding bypass. When False
(default), omitting ``session_id`` is acceptable even if the token
contains one.
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
token_session_id = payload.get("session_id")
if require_session and token_session_id is not None:
if session_id is None or session_id != token_session_id:
return None
elif session_id is not None and session_id != token_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))
if random.random() < 0.02:
blacklist_expired()
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:
"""Sliding-window rate limiter that tracks successes and failures separately.
Failures are counted against the limit. A successful operation resets
the failure counter for that key.
"""
def __init__(self, max_attempts: int = 5, window_seconds: int = 300) -> None:
self.max_attempts = max_attempts
self.window = window_seconds
self.failures: dict[str, list[float]] = {}
def is_allowed(self, key: str) -> bool:
"""Check if a request from *key* is allowed (does NOT record the attempt).
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.failures.get(key, [])
clean = [t for t in timestamps if t > cutoff]
return len(clean) < self.max_attempts
def record_failure(self, key: str) -> None:
"""Record a failed attempt for *key*."""
self.failures.setdefault(key, []).append(time.time())
def record_success(self, key: str) -> None:
"""Reset the failure counter for *key* on a successful operation."""
self.failures.pop(key, None)
def cleanup(self) -> None:
"""Remove expired entries from all buckets."""
now = time.time()
cutoff = now - self.window
for key in list(self.failures):
self.failures[key] = [t for t in self.failures[key] if t > cutoff]
if not self.failures[key]:
del self.failures[key]
# Global rate limiters — in-memory only. Counts reset on daemon restart
# (SIGHUP reload, process restart). Acceptable for a single-user appliance
# where restarts are rare; brute-force windows briefly reset post-restart.
_login_limiter = RateLimiter(max_attempts=10, window_seconds=300)
_webauthn_limiter = RateLimiter(max_attempts=5, window_seconds=600)
def check_login_rate(username: str, client_ip: str | None = None) -> bool:
"""Check if login is rate-limited for the given username.
Checks failure counts for both IP and username buckets without
recording anything. Callers must invoke record_login_failure() or
record_login_success() after the password verification step.
Args:
username: The login attempt username.
client_ip: The client IP address (from X-Real-IP header).
Returns:
True if the attempt is allowed, False if rate limited.
"""
if client_ip and not _login_limiter.is_allowed(client_ip):
return False
return _login_limiter.is_allowed(username)
def record_login_failure(username: str, client_ip: str | None = None) -> None:
"""Record a failed login attempt."""
if client_ip:
_login_limiter.record_failure(client_ip)
_login_limiter.record_failure(username)
def record_login_success(username: str, client_ip: str | None = None) -> None:
"""Record a successful login (resets failure counter)."""
if client_ip:
_login_limiter.record_success(client_ip)
_login_limiter.record_success(username)
def check_webauthn_rate(username: str, client_ip: str | None = None) -> bool:
"""Check if WebAuthn authentication is rate-limited for the given username.
Args:
username: The WebAuthn attempt username.
client_ip: The client IP address (from X-Real-IP header).
Returns:
True if the attempt is allowed, False if rate limited.
"""
if client_ip and not _webauthn_limiter.is_allowed(client_ip):
return False
return _webauthn_limiter.is_allowed(username)
def record_webauthn_failure(username: str, client_ip: str | None = None) -> None:
"""Record a failed WebAuthn attempt."""
if client_ip:
_webauthn_limiter.record_failure(client_ip)
_webauthn_limiter.record_failure(username)
def record_webauthn_success(username: str, client_ip: str | None = None) -> None:
"""Record a successful WebAuthn attempt (resets failure counter)."""
if client_ip:
_webauthn_limiter.record_success(client_ip)
_webauthn_limiter.record_success(username)