306 lines
11 KiB
Python
306 lines
11 KiB
Python
# backend/auth/jwt_handler.py
|
|
# JWT token generation, validation, and revocation.
|
|
# Secret key auto-generated on first startup and persisted to data/secret.key.
|
|
# Revoked token JTIs persisted to data/revoked_tokens.json.
|
|
|
|
import json
|
|
import logging
|
|
import os
|
|
import secrets
|
|
import threading
|
|
import time
|
|
import uuid
|
|
from pathlib import Path
|
|
|
|
from jose import JWTError, jwt
|
|
|
|
logger = logging.getLogger("obsigate.auth.jwt")
|
|
|
|
# Paths relative to working directory (Docker: /app)
|
|
SECRET_KEY_FILE = Path("data/secret.key")
|
|
REVOKED_TOKENS_FILE = Path("data/revoked_tokens.json")
|
|
|
|
ALGORITHM = "HS256"
|
|
ACCESS_TOKEN_EXPIRE_SECONDS = int(os.environ.get("OBSIGATE_ACCESS_TOKEN_TTL", "3600")) # default 1 hour
|
|
REFRESH_TOKEN_EXPIRE_SECONDS = int(os.environ.get("OBSIGATE_REFRESH_TOKEN_TTL", "604800")) # default 7 days
|
|
|
|
#: Persistent API/MCP access tokens (user-managed, shown in the config panel).
|
|
API_TOKENS_FILE = Path("data/api_tokens.json")
|
|
#: Accepted values for the expiry selector in the UI (1 day, 1 month, 6 months,
|
|
#: 1 year, never). "never" → no ``exp`` claim → token valid until revoked.
|
|
API_TOKEN_EXPIRY_CHOICES = {
|
|
"1d": 24 * 3600,
|
|
"30d": 30 * 24 * 3600,
|
|
"180d": 180 * 24 * 3600,
|
|
"365d": 365 * 24 * 3600,
|
|
"never": None,
|
|
}
|
|
#: Max active tokens per user (anti hoarding; revoking frees a slot).
|
|
API_TOKEN_MAX_PER_USER = 50
|
|
#: AES-GCM key derived once from the JWT secret to encrypt stored tokens.
|
|
_API_TOKEN_KEY: bytes | None = None
|
|
|
|
# In-memory revoked token set (loaded from disk on startup)
|
|
_revoked_jtis: set = set()
|
|
_revoked_loaded = False
|
|
|
|
|
|
def get_secret_key() -> str:
|
|
"""Read or generate the JWT secret key.
|
|
|
|
On first call, generates a 512-bit random key and writes it to
|
|
data/secret.key with 600 permissions. Subsequent calls read from disk.
|
|
"""
|
|
if not SECRET_KEY_FILE.exists():
|
|
SECRET_KEY_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
key = secrets.token_hex(64) # 512 bits
|
|
SECRET_KEY_FILE.write_text(key)
|
|
try:
|
|
SECRET_KEY_FILE.chmod(0o600)
|
|
except OSError:
|
|
pass # Windows doesn't support Unix permissions
|
|
logger.info("Generated new JWT secret key")
|
|
return key
|
|
return SECRET_KEY_FILE.read_text().strip()
|
|
|
|
|
|
def create_access_token(user: dict) -> str:
|
|
"""Create a JWT access token with user claims."""
|
|
now = int(time.time())
|
|
payload = {
|
|
"sub": user["username"],
|
|
"role": user["role"],
|
|
"vaults": user["vaults"],
|
|
"jti": str(uuid.uuid4()),
|
|
"iat": now,
|
|
"exp": now + ACCESS_TOKEN_EXPIRE_SECONDS,
|
|
"type": "access",
|
|
}
|
|
return jwt.encode(payload, get_secret_key(), algorithm=ALGORITHM)
|
|
|
|
|
|
def create_refresh_token(username: str, remember: bool = False) -> tuple:
|
|
"""Create a JWT refresh token. Returns (token_string, jti).
|
|
|
|
``remember`` is carried as a claim so token rotation can preserve the
|
|
30-day vs 7-day lifetime chosen at login.
|
|
"""
|
|
now = int(time.time())
|
|
jti = str(uuid.uuid4())
|
|
payload = {
|
|
"sub": username,
|
|
"jti": jti,
|
|
"iat": now,
|
|
"exp": now + (2592000 if remember else REFRESH_TOKEN_EXPIRE_SECONDS),
|
|
"type": "refresh",
|
|
"remember": remember,
|
|
}
|
|
return jwt.encode(payload, get_secret_key(), algorithm=ALGORITHM), jti
|
|
|
|
|
|
def decode_token(token: str) -> dict | None:
|
|
"""Decode and validate a JWT. Returns None if invalid/expired."""
|
|
try:
|
|
return jwt.decode(token, get_secret_key(), algorithms=[ALGORITHM])
|
|
except JWTError:
|
|
return None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Token revocation
|
|
# ---------------------------------------------------------------------------
|
|
# The store is a dict {jti: valid_until}: the revocation record may be dropped
|
|
# once the underlying token's own expiry has passed (by then the JWT is dead
|
|
# anyway). Long-lived API/MCP tokens (see create_api_token) must therefore be
|
|
# revoked with their real expiry — a 1-year token revoked last week must not
|
|
# silently come back to life when a 7-day cleanup purges the record (BUG in
|
|
# the previous set-based store, fixed with feature #107).
|
|
|
|
_revoked_map: dict[str, int] = {}
|
|
_revoked_loaded = False
|
|
|
|
# ROADMAP #85 T10a — verrou autour du read-modify-write du store de
|
|
# révocation (perte de révocations en cas de logouts concurrents).
|
|
_revoked_lock = threading.RLock()
|
|
|
|
|
|
def _load_revoked():
|
|
"""Load revoked token JTIs from disk into memory (once)."""
|
|
global _revoked_loaded, _revoked_map
|
|
with _revoked_lock:
|
|
if _revoked_loaded:
|
|
return
|
|
if REVOKED_TOKENS_FILE.exists():
|
|
try:
|
|
data = json.loads(REVOKED_TOKENS_FILE.read_text())
|
|
# Drop entries whose underlying token has itself expired.
|
|
now = int(time.time())
|
|
_revoked_map = {
|
|
jti: int(exp) for jti, exp in data.items()
|
|
if int(exp) > now
|
|
}
|
|
except Exception as e:
|
|
logger.warning(f"Failed to load revoked tokens: {e}")
|
|
_revoked_map = {}
|
|
_revoked_loaded = True
|
|
|
|
|
|
def _save_revoked():
|
|
"""Persist revoked JTIs to disk with their per-token expiry."""
|
|
REVOKED_TOKENS_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
tmp = REVOKED_TOKENS_FILE.with_suffix(".tmp")
|
|
tmp.write_text(json.dumps(_revoked_map))
|
|
tmp.replace(REVOKED_TOKENS_FILE)
|
|
|
|
|
|
def revoke_token(jti: str, expires_at: int | None = None):
|
|
"""Add a token JTI to the revocation list.
|
|
|
|
``expires_at`` is the revoked token's own ``exp`` (unix seconds) — the
|
|
record is kept at least that long so a long-lived API token cannot
|
|
outlive its revocation. ``None`` means the token never expires (API/MCP
|
|
"sans fin") → the record is kept forever (capped at ~100 years, the JWT
|
|
store's practical infinity). Default keeps 7 days (session tokens).
|
|
"""
|
|
with _revoked_lock:
|
|
_load_revoked()
|
|
now = int(time.time())
|
|
if expires_at is None:
|
|
until = now + 100 * 365 * 24 * 3600
|
|
else:
|
|
until = max(int(expires_at), now + REFRESH_TOKEN_EXPIRE_SECONDS)
|
|
_revoked_map[jti] = until
|
|
_save_revoked()
|
|
logger.debug(f"Revoked token JTI: {jti[:8]}...")
|
|
|
|
|
|
def is_token_revoked(jti: str) -> bool:
|
|
"""Check if a token JTI has been revoked."""
|
|
with _revoked_lock:
|
|
_load_revoked()
|
|
return jti in _revoked_map
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# API / MCP tokens (feature #107)
|
|
# ---------------------------------------------------------------------------
|
|
# Long-lived access tokens the user creates from the config panel. They are
|
|
# plain HS256 access-type JWTs (``api: true`` claim), so they authenticate
|
|
# against BOTH the REST API and the MCP endpoint (/mcp) — which share
|
|
# ``get_current_user``. The raw token is shown exactly once at creation; the
|
|
# store keeps metadata only (name, owner, expiry, last use) — no secret
|
|
# material is written to disk.
|
|
#
|
|
# File: data/api_tokens.json
|
|
# {"version": 1, "tokens": {jti: {name, username, created_at, expires_at, last_used_at}}}
|
|
|
|
_api_tokens_lock = threading.RLock()
|
|
_touch_last_write: dict[str, float] = {}
|
|
|
|
|
|
def _load_api_tokens() -> dict:
|
|
if not API_TOKENS_FILE.exists():
|
|
return {"version": 1, "tokens": {}}
|
|
try:
|
|
return json.loads(API_TOKENS_FILE.read_text(encoding="utf-8"))
|
|
except (json.JSONDecodeError, OSError) as e:
|
|
logger.error(f"Failed to read api_tokens.json: {e}")
|
|
return {"version": 1, "tokens": {}}
|
|
|
|
|
|
def _save_api_tokens(data: dict):
|
|
API_TOKENS_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
tmp = API_TOKENS_FILE.with_suffix(".tmp")
|
|
tmp.write_text(json.dumps(data, indent=2, default=str), encoding="utf-8")
|
|
tmp.replace(API_TOKENS_FILE)
|
|
|
|
|
|
def create_api_token(user: dict, name: str, expiry_key: str) -> tuple[dict, str]:
|
|
"""Create a persistent API/MCP token. Returns (record, jwt_string).
|
|
|
|
``expiry_key`` must be one of API_TOKEN_EXPIRY_CHOICES; "never" omits the
|
|
``exp`` claim (valid until explicitly revoked).
|
|
"""
|
|
if expiry_key not in API_TOKEN_EXPIRY_CHOICES:
|
|
raise ValueError("Expiration invalide")
|
|
seconds = API_TOKEN_EXPIRY_CHOICES[expiry_key]
|
|
with _api_tokens_lock:
|
|
data = _load_api_tokens()
|
|
tokens = data["tokens"]
|
|
mine = sum(1 for t in tokens.values() if t["username"] == user["username"])
|
|
if mine >= API_TOKEN_MAX_PER_USER:
|
|
raise ValueError(f"Maximum {API_TOKEN_MAX_PER_USER} tokens par utilisateur")
|
|
now = int(time.time())
|
|
jti = str(uuid.uuid4())
|
|
payload = {
|
|
"sub": user["username"],
|
|
"role": user.get("role", "user"),
|
|
"vaults": user.get("vaults", []),
|
|
"jti": jti,
|
|
"iat": now,
|
|
"type": "access",
|
|
"api": True,
|
|
}
|
|
if seconds is not None:
|
|
payload["exp"] = now + seconds
|
|
token = jwt.encode(payload, get_secret_key(), algorithm=ALGORITHM)
|
|
record = {
|
|
"jti": jti,
|
|
"name": name[:64] or "API token",
|
|
"username": user["username"],
|
|
"created_at": now,
|
|
"expires_at": payload.get("exp"),
|
|
"expiry_key": expiry_key,
|
|
"last_used_at": None,
|
|
}
|
|
tokens[jti] = record
|
|
_save_api_tokens(data)
|
|
return record, token
|
|
|
|
|
|
def list_api_tokens(username: str) -> list[dict]:
|
|
"""Token metadata for one user, newest first."""
|
|
data = _load_api_tokens()
|
|
now = int(time.time())
|
|
items = [
|
|
{**t, "expired": t.get("expires_at") is not None and t["expires_at"] < now}
|
|
for t in data["tokens"].values()
|
|
if t["username"] == username
|
|
]
|
|
return sorted(items, key=lambda t: t["created_at"], reverse=True)
|
|
|
|
|
|
def delete_api_token(jti: str, username: str) -> dict:
|
|
"""Revoke and remove an API token. Raises KeyError when unknown/not owned."""
|
|
with _api_tokens_lock:
|
|
data = _load_api_tokens()
|
|
record = data["tokens"].get(jti)
|
|
if not record or record["username"] != username:
|
|
raise KeyError(jti)
|
|
# Revoke by jti so the presented JWT stops working even though it is
|
|
# stateless — kept until its natural expiry (no-expiry → forever).
|
|
revoke_token(jti, record.get("expires_at"))
|
|
del data["tokens"][jti]
|
|
_save_api_tokens(data)
|
|
return record
|
|
|
|
|
|
def maybe_touch_api_token(jti: str | None, created_or_expires: bool = False):
|
|
"""Record last usage of an API token, throttled to one disk write/hour."""
|
|
if not jti:
|
|
return
|
|
now = time.time()
|
|
if now - _touch_last_write.get(jti, 0) < 3600:
|
|
return
|
|
_touch_last_write[jti] = now
|
|
try:
|
|
with _api_tokens_lock:
|
|
data = _load_api_tokens()
|
|
record = data["tokens"].get(jti)
|
|
if record is None:
|
|
return
|
|
record["last_used_at"] = int(now)
|
|
_save_api_tokens(data)
|
|
except Exception as e: # never fail an authenticated request over stats
|
|
logger.debug(f"api_token touch failed: {e}")
|