521 lines
17 KiB
Python
521 lines
17 KiB
Python
"""
|
|
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)
|