""" auth.py — Shared authentication & encryption module. Responsibilities: - Derive or generate a master secret key (persisted to disk with chmod 600) - Provide Fernet-based encryption for user PII at rest - Argon2id password hashing (constant-time verification) - TOTP secret generation, QR provisioning URI, 6-digit code verification - JWT session tokens (signed with master key) - User CRUD on the shared SQLite database - FastAPI dependencies to enforce authentication and admin role """ from __future__ import annotations import base64 import hashlib import hmac import os import secrets import sqlite3 import time from contextlib import contextmanager from dataclasses import dataclass from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Optional import jwt import pyotp from argon2 import PasswordHasher from argon2.exceptions import VerifyMismatchError, InvalidHashError from cryptography.fernet import Fernet, InvalidToken from cryptography.hazmat.primitives import hashes from cryptography.hazmat.primitives.kdf.hkdf import HKDF from fastapi import Depends, HTTPException, Request, status # --------------------------------------------------------------------------- # Config # --------------------------------------------------------------------------- DATA_DIR = Path(os.environ.get("DATA_DIR", "/app/data")) DATA_DIR.mkdir(parents=True, exist_ok=True) SECRET_KEY_PATH = DATA_DIR / "secret.key" DB_PATH = DATA_DIR / "pst_index.db" SESSION_TTL_SECONDS = int(os.environ.get("SESSION_TTL_SECONDS", 60 * 60 * 8)) # 8 hours SESSION_COOKIE_NAME = "pst_session" ADMIN_SESSION_COOKIE_NAME = "pst_admin_session" JWT_ALGORITHM = "HS256" # --------------------------------------------------------------------------- # Master secret # --------------------------------------------------------------------------- def _load_or_create_master_secret() -> bytes: """Return a 32-byte master secret. Created on first run, persisted to disk.""" # Allow overriding via env var (useful for orchestration); still persist so # restarts without the env var keep working. env_secret = os.environ.get("MASTER_SECRET") if env_secret: return hashlib.sha256(env_secret.encode("utf-8")).digest() if SECRET_KEY_PATH.exists(): data = SECRET_KEY_PATH.read_bytes() if len(data) == 32: return data # Legacy or truncated: rewrap return hashlib.sha256(data).digest() secret = secrets.token_bytes(32) SECRET_KEY_PATH.write_bytes(secret) try: os.chmod(SECRET_KEY_PATH, 0o600) except Exception: pass return secret MASTER_SECRET: bytes = _load_or_create_master_secret() def _derive_subkey(info: bytes, length: int = 32) -> bytes: """Derive a sub-key from the master secret using HKDF-SHA256.""" return HKDF( algorithm=hashes.SHA256(), length=length, salt=b"pst-indexer-salt-v1", info=info, ).derive(MASTER_SECRET) # Key used for encrypting user PII at rest _FERNET_KEY = base64.urlsafe_b64encode(_derive_subkey(b"fernet-user-pii")) _FERNET = Fernet(_FERNET_KEY) # Key used for JWT signatures _JWT_KEY = _derive_subkey(b"jwt-sessions") # Key used for deterministic username lookup hashes (so we can index without # exposing plaintext usernames in the DB). _USERNAME_HMAC_KEY = _derive_subkey(b"username-lookup") # --------------------------------------------------------------------------- # Encryption helpers # --------------------------------------------------------------------------- def encrypt(plaintext: str) -> str: """Fernet-encrypt a string; returns base64 ciphertext.""" if plaintext is None: return "" return _FERNET.encrypt(plaintext.encode("utf-8")).decode("utf-8") def decrypt(ciphertext: str) -> str: """Fernet-decrypt. Returns empty string for empty input.""" if not ciphertext: return "" try: return _FERNET.decrypt(ciphertext.encode("utf-8")).decode("utf-8") except InvalidToken: raise ValueError("Failed to decrypt — master key may have changed") def username_lookup_hash(username: str) -> str: """Deterministic HMAC-SHA256 over the normalized username, hex-encoded.""" normalized = username.strip().lower().encode("utf-8") return hmac.new(_USERNAME_HMAC_KEY, normalized, hashlib.sha256).hexdigest() # --------------------------------------------------------------------------- # Password hashing # --------------------------------------------------------------------------- _hasher = PasswordHasher( time_cost=3, memory_cost=64 * 1024, # 64 MiB parallelism=2, ) def hash_password(password: str) -> str: return _hasher.hash(password) def verify_password(stored_hash: str, password: str) -> bool: try: _hasher.verify(stored_hash, password) return True except (VerifyMismatchError, InvalidHashError): return False except Exception: return False # --------------------------------------------------------------------------- # TOTP (MFA) # --------------------------------------------------------------------------- def generate_totp_secret() -> str: return pyotp.random_base32() def totp_provisioning_uri(secret: str, username: str, issuer: str = "PST Archive") -> str: return pyotp.TOTP(secret).provisioning_uri(name=username, issuer_name=issuer) def verify_totp(secret: str, code: str) -> bool: if not secret or not code: return False try: return pyotp.TOTP(secret).verify(code.strip(), valid_window=1) except Exception: return False # --------------------------------------------------------------------------- # Database # --------------------------------------------------------------------------- @contextmanager def get_db(): conn = sqlite3.connect(str(DB_PATH)) conn.row_factory = sqlite3.Row conn.execute("PRAGMA foreign_keys = ON") try: yield conn finally: conn.close() def init_auth_schema(): with get_db() as conn: conn.executescript(""" CREATE TABLE IF NOT EXISTS users ( id TEXT PRIMARY KEY, username_hash TEXT NOT NULL UNIQUE, username_enc TEXT NOT NULL, password_hash_enc TEXT NOT NULL, totp_secret_enc TEXT, mfa_enabled INTEGER NOT NULL DEFAULT 0, role TEXT NOT NULL DEFAULT 'user', created_at TEXT NOT NULL, last_login TEXT ); CREATE INDEX IF NOT EXISTS idx_users_username_hash ON users(username_hash); CREATE INDEX IF NOT EXISTS idx_users_role ON users(role); """) conn.commit() # --------------------------------------------------------------------------- # User model # --------------------------------------------------------------------------- @dataclass class User: id: str username: str role: str mfa_enabled: bool created_at: str last_login: Optional[str] @property def is_admin(self) -> bool: return self.role == "admin" def public_dict(self) -> dict: return { "id": self.id, "username": self.username, "role": self.role, "mfa_enabled": self.mfa_enabled, "created_at": self.created_at, "last_login": self.last_login, } def _row_to_user(row: sqlite3.Row) -> User: return User( id=row["id"], username=decrypt(row["username_enc"]), role=row["role"], mfa_enabled=bool(row["mfa_enabled"]), created_at=row["created_at"], last_login=row["last_login"], ) def count_users() -> int: with get_db() as conn: return conn.execute("SELECT COUNT(*) AS c FROM users").fetchone()["c"] def count_admins() -> int: with get_db() as conn: return conn.execute("SELECT COUNT(*) AS c FROM users WHERE role='admin'").fetchone()["c"] def find_user_by_username(username: str) -> Optional[tuple[User, str]]: """Return (User, decrypted_password_hash) if found, else None.""" uh = username_lookup_hash(username) with get_db() as conn: row = conn.execute( "SELECT * FROM users WHERE username_hash = ?", (uh,) ).fetchone() if not row: return None user = _row_to_user(row) try: pw_hash = decrypt(row["password_hash_enc"]) except ValueError: return None return user, pw_hash def find_user_by_id(user_id: str) -> Optional[User]: with get_db() as conn: row = conn.execute("SELECT * FROM users WHERE id = ?", (user_id,)).fetchone() if not row: return None return _row_to_user(row) def get_totp_secret(user_id: str) -> Optional[str]: with get_db() as conn: row = conn.execute( "SELECT totp_secret_enc FROM users WHERE id = ?", (user_id,) ).fetchone() if not row or not row["totp_secret_enc"]: return None try: return decrypt(row["totp_secret_enc"]) except ValueError: return None def list_users() -> list[User]: with get_db() as conn: rows = conn.execute( "SELECT * FROM users ORDER BY created_at ASC" ).fetchall() return [_row_to_user(r) for r in rows] def create_user(username: str, password: str, role: str = "user") -> User: username = username.strip() if not username or len(username) < 3: raise ValueError("Username must be at least 3 characters") if len(username) > 64: raise ValueError("Username too long") if role not in ("user", "admin"): raise ValueError("Invalid role") if len(password) < 8: raise ValueError("Password must be at least 8 characters") uh = username_lookup_hash(username) with get_db() as conn: existing = conn.execute( "SELECT id FROM users WHERE username_hash = ?", (uh,) ).fetchone() if existing: raise ValueError("Username already taken") user_id = secrets.token_hex(16) now = datetime.now(timezone.utc).isoformat() pw_hash = hash_password(password) conn.execute( """INSERT INTO users (id, username_hash, username_enc, password_hash_enc, totp_secret_enc, mfa_enabled, role, created_at, last_login) VALUES (?, ?, ?, ?, NULL, 0, ?, ?, NULL)""", (user_id, uh, encrypt(username), encrypt(pw_hash), role, now), ) conn.commit() return User( id=user_id, username=username, role=role, mfa_enabled=False, created_at=now, last_login=None, ) def delete_user(user_id: str) -> None: with get_db() as conn: # Also clean up any folder permission grants this user had conn.execute("DELETE FROM folder_permissions WHERE user_id = ?", (user_id,)) conn.execute("DELETE FROM users WHERE id = ?", (user_id,)) conn.commit() def set_user_role(user_id: str, role: str) -> None: if role not in ("user", "admin"): raise ValueError("Invalid role") with get_db() as conn: conn.execute("UPDATE users SET role = ? WHERE id = ?", (role, user_id)) conn.commit() def update_password(user_id: str, new_password: str) -> None: if len(new_password) < 8: raise ValueError("Password must be at least 8 characters") pw_hash = hash_password(new_password) with get_db() as conn: conn.execute( "UPDATE users SET password_hash_enc = ? WHERE id = ?", (encrypt(pw_hash), user_id), ) conn.commit() def begin_mfa_enrollment(user_id: str) -> tuple[str, str]: """Generate a new TOTP secret (not yet enabled) and return (secret, provisioning_uri).""" user = find_user_by_id(user_id) if not user: raise ValueError("User not found") secret = generate_totp_secret() # Stash it but leave mfa_enabled = 0 until confirmed with get_db() as conn: conn.execute( "UPDATE users SET totp_secret_enc = ?, mfa_enabled = 0 WHERE id = ?", (encrypt(secret), user_id), ) conn.commit() uri = totp_provisioning_uri(secret, user.username) return secret, uri def confirm_mfa_enrollment(user_id: str, code: str) -> bool: secret = get_totp_secret(user_id) if not secret or not verify_totp(secret, code): return False with get_db() as conn: conn.execute("UPDATE users SET mfa_enabled = 1 WHERE id = ?", (user_id,)) conn.commit() return True def disable_mfa(user_id: str) -> None: with get_db() as conn: conn.execute( "UPDATE users SET mfa_enabled = 0, totp_secret_enc = NULL WHERE id = ?", (user_id,), ) conn.commit() def record_login(user_id: str) -> None: with get_db() as conn: conn.execute( "UPDATE users SET last_login = ? WHERE id = ?", (datetime.now(timezone.utc).isoformat(), user_id), ) conn.commit() # --------------------------------------------------------------------------- # JWT sessions # --------------------------------------------------------------------------- def create_session_token(user: User, scope: str = "app") -> str: """scope='app' for main app, scope='admin' for admin panel.""" now = datetime.now(timezone.utc) payload = { "sub": user.id, "role": user.role, "scope": scope, "iat": int(now.timestamp()), "exp": int((now + timedelta(seconds=SESSION_TTL_SECONDS)).timestamp()), } return jwt.encode(payload, _JWT_KEY, algorithm=JWT_ALGORITHM) def verify_session_token(token: str, expected_scope: str) -> Optional[dict]: try: payload = jwt.decode(token, _JWT_KEY, algorithms=[JWT_ALGORITHM]) if payload.get("scope") != expected_scope: return None return payload except jwt.PyJWTError: return None # --------------------------------------------------------------------------- # FastAPI dependencies # --------------------------------------------------------------------------- def _extract_token(request: Request, cookie_name: str) -> Optional[str]: # Prefer cookie, fall back to Authorization header token = request.cookies.get(cookie_name) if token: return token auth = request.headers.get("authorization", "") if auth.lower().startswith("bearer "): return auth[7:].strip() return None def require_user(scope: str, cookie_name: str): def _dep(request: Request) -> User: token = _extract_token(request, cookie_name) if not token: raise HTTPException(status.HTTP_401_UNAUTHORIZED, "Not authenticated") payload = verify_session_token(token, expected_scope=scope) if not payload: raise HTTPException(status.HTTP_401_UNAUTHORIZED, "Invalid or expired session") user = find_user_by_id(payload["sub"]) if not user: raise HTTPException(status.HTTP_401_UNAUTHORIZED, "User no longer exists") return user return _dep def require_admin(scope: str, cookie_name: str): base = require_user(scope, cookie_name) def _dep(request: Request) -> User: user = base(request) if not user.is_admin: raise HTTPException(status.HTTP_403_FORBIDDEN, "Admin role required") return user return _dep # Pre-wired dependencies for the two app scopes require_app_user = require_user("app", SESSION_COOKIE_NAME) require_admin_user = require_admin("admin", ADMIN_SESSION_COOKIE_NAME) # App-scoped admin dependency: authenticates against the MAIN app session # cookie (scope "app") but additionally requires the admin role. This lets the # administration section run inside the main application (served at /admin) # instead of as a separate process on its own port, while still being strictly # limited to admin users. require_app_admin = require_admin("app", SESSION_COOKIE_NAME) # --------------------------------------------------------------------------- # Simple rate limiter (per-IP, per-username) to slow brute-force attempts # --------------------------------------------------------------------------- class RateLimiter: def __init__(self, window_seconds: int = 300, max_attempts: int = 10): self.window = window_seconds self.max = max_attempts self._attempts: dict[str, list[float]] = {} def check(self, key: str) -> bool: now = time.time() cutoff = now - self.window arr = [t for t in self._attempts.get(key, []) if t > cutoff] if len(arr) >= self.max: self._attempts[key] = arr return False return True def record(self, key: str) -> None: now = time.time() cutoff = now - self.window arr = [t for t in self._attempts.get(key, []) if t > cutoff] arr.append(now) self._attempts[key] = arr def reset(self, key: str) -> None: self._attempts.pop(key, None) login_limiter = RateLimiter(window_seconds=300, max_attempts=10)