"""Magic-link auth, sessions, staff-key, rate limiting (in-memory sliding window). Constant-time comparison via hmac.compare_digest for token hashes + staff key. """ from __future__ import annotations import hashlib import hmac import json import logging import os import secrets import time import urllib.error import urllib.parse import urllib.request from collections import defaultdict, deque from datetime import datetime, timedelta, timezone from fastapi import Depends, HTTPException, Request, status from . import db from .db import get_conn, new_uuid, utcnow_iso logger = logging.getLogger("dre.auth") MAGIC_TOKEN_TTL_MIN = 15 SESSION_TTL_DAYS = 7 # --------------------------------------------------------------- # Rate limiting — in-memory sliding window (single-instance v1) # --------------------------------------------------------------- class RateLimiter: def __init__(self) -> None: self._by_email: dict[str, deque[float]] = defaultdict(deque) self._by_ip: dict[str, deque[float]] = defaultdict(deque) def _prune(self, dq: deque[float], window_sec: float) -> None: cutoff = time.time() - window_sec while dq and dq[0] < cutoff: dq.popleft() def check_email(self, email: str, max_count: int, window_sec: float) -> bool: dq = self._by_email[email] self._prune(dq, window_sec) if len(dq) >= max_count: return False dq.append(time.time()) return True def check_ip(self, ip: str, max_count: int, window_sec: float) -> bool: dq = self._by_ip[ip] self._prune(dq, window_sec) if len(dq) >= max_count: return False dq.append(time.time()) return True def check(self, email: str | None, ip: str, max_email: int, email_window: float, max_ip: int, ip_window: float) -> bool: if email and not self.check_email(email, max_email, email_window): return False if not self.check_ip(ip, max_ip, ip_window): return False return True _limiter = RateLimiter() def rate_limit_auth_request(email: str | None, ip: str) -> bool: return _limiter.check(email, ip, max_email=3, email_window=900, max_ip=10, ip_window=3600) def rate_limit_auth_verify(ip: str) -> bool: return _limiter.check(None, ip, max_email=999, email_window=1, max_ip=10, ip_window=900) def rate_limit_intake(ip: str) -> bool: return _limiter.check_ip(ip, max_count=20, window_sec=3600) # --------------------------------------------------------------- # Token / hash helpers # --------------------------------------------------------------- def _sha256_hex(value: str) -> str: return hashlib.sha256(value.encode("utf-8")).hexdigest() def generate_magic_token() -> str: return secrets.token_urlsafe(32) def generate_session_token() -> str: return secrets.token_urlsafe(32) def _compare_hash(a: str, b: str) -> bool: return hmac.compare_digest(a, b) # --------------------------------------------------------------- # Staff key # --------------------------------------------------------------- def _staff_key() -> str: return os.environ.get("DRE_STAFF_KEY", "") def verify_staff_key(provided: str | None) -> bool: key = _staff_key() if not key or not provided: return False return _compare_hash(provided, key) # --------------------------------------------------------------- # Stack Auth (auth2 / Hexclave) — per-user staff SSO # --------------------------------------------------------------- STACK_AUTH_API_URL = os.environ.get("STACK_AUTH_API_URL", "https://auth2-api.itpropartner.com").rstrip("/") STACK_AUTH_PUBLISHABLE_KEY = os.environ.get("STACK_AUTH_PUBLISHABLE_KEY", "") STACK_AUTH_PROJECT_ID = os.environ.get("STACK_AUTH_PROJECT_ID", "internal") STACK_AUTH_TEAM = os.environ.get("STACK_AUTH_TEAM", "dre-staff") # Emergency owner bypass: these emails are always admitted even if the team # membership API is down or membership was accidentally removed. STACK_AUTH_OWNER_EMAILS = [ e.strip().lower() for e in os.environ.get("STACK_AUTH_OWNER_EMAILS", "").split(",") if e.strip() ] def _stack_auth_request(method: str, path: str, access_token: str | None = None, body: dict | None = None): """Call the auth2 Stack Auth REST API (client mode). Returns (status, parsed|raw). Unauthenticated calls use the publishable client key; authenticated calls use the opaque access token instead. Both carry access-type/project-id headers. """ url = STACK_AUTH_API_URL + path headers = { "Content-Type": "application/json", "x-hexclave-access-type": "client", "x-hexclave-project-id": STACK_AUTH_PROJECT_ID, } if access_token: headers["x-hexclave-access-token"] = access_token else: headers["x-hexclave-publishable-client-key"] = STACK_AUTH_PUBLISHABLE_KEY data = json.dumps(body).encode("utf-8") if body is not None else None req = urllib.request.Request(url, data=data, method=method, headers=headers) try: resp = urllib.request.urlopen(req, timeout=8) raw = resp.read().decode("utf-8", "replace") code = resp.status except urllib.error.HTTPError as e: code = e.code raw = e.read().decode("utf-8", "replace") except Exception: # noqa: BLE001 return None, None try: return code, json.loads(raw) except Exception: # noqa: BLE001 return code, raw def stack_auth_user(access_token: str | None) -> tuple[str, dict | None]: """Validate an opaque Stack Auth access token and confirm dre-staff membership. Returns one of: ("ok", {"name","email","user_id"}) — valid session AND team member ("denied", None) — valid session but NOT in the team ("invalid", None) — bad/expired token or API unreachable """ if not access_token or not STACK_AUTH_PUBLISHABLE_KEY: return "invalid", None code, user = _stack_auth_request("GET", "/api/latest/users/me", access_token=access_token) if code != 200 or not isinstance(user, dict): return "invalid", None q = urllib.parse.quote(STACK_AUTH_TEAM) code, teams = _stack_auth_request("GET", f"/api/latest/teams?user_id=me&query={q}", access_token=access_token) items = teams.get("items") if isinstance(teams, dict) else None member = (code == 200 and isinstance(items, list) and any( str(t.get("display_name", "")).lower() == STACK_AUTH_TEAM.lower() for t in items)) email = user.get("primary_email") or "" is_owner = email.lower() in STACK_AUTH_OWNER_EMAILS if not member and not is_owner: return "denied", None return "ok", { "name": user.get("display_name") or email or "DRE Staff", "email": email, "user_id": user.get("id") or "", } def stack_auth_sign_in(email: str, password: str) -> tuple[str, dict | None]: """Password sign-in against Stack Auth, gated to the dre-staff team. Returns ("ok", {"token","name","email","user_id"}), ("invalid", None) for bad credentials, or ("denied", None) for a valid account outside the team. """ code, d = _stack_auth_request( "POST", "/api/latest/auth/password/sign-in", body={"email": email, "password": password}, ) if code != 200 or not isinstance(d, dict): return "invalid", None token = d.get("access_token") if not token: return "invalid", None state, user = stack_auth_user(token) if state != "ok": return "denied", None return "ok", {"token": token, **user} # --------------------------------------------------------------- # FastAPI dependencies # --------------------------------------------------------------- def get_client_ip(request: Request) -> str: # Cloudflare / Caddy may set X-Forwarded-For; use first hop xff = request.headers.get("x-forwarded-for") if xff: return xff.split(",")[0].strip() return request.client.host if request.client else "unknown" def staff_identity(provided: str | None) -> str: """Resolve the acting staff member's display name from the presented key. Priority: 1. DRE_STAFF_DIRECTORY — 'KEY=Name ' entries (one per line or semicolon-separated). Enables real per-user RBAC once each staff member has their own key. 2. DRE_STAFF_NAME — a single default name for the shared-key case. 3. "DRE Staff" fallback. """ name = os.environ.get("DRE_STAFF_NAME", "").strip() directory = os.environ.get("DRE_STAFF_DIRECTORY", "").strip() if directory and provided: for entry in directory.replace(";", "\n").splitlines(): entry = entry.strip() if not entry or "=" not in entry: continue key, _, label = entry.partition("=") if key.strip() == provided: name = label.strip() break return name or "DRE Staff" def require_staff(request: Request) -> str: """Staff auth dependency. Raises 403 if missing/wrong. Priority: 1. Stack Auth access token (`x-dre-access-token`) → per-user auth2 SSO, validated server-side and gated to the dre-staff team. 2. Legacy staff key (`x-dre-staff-key`) → script/fallback access. Returns the acting staff member's display name so callers can record the actor without manual name entry. """ token = request.headers.get("x-dre-access-token") if token: state, user = stack_auth_user(token) if state == "ok": return user["name"] if state == "denied": raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={"code": "staff_forbidden", "message": "Account is not authorized for staff access."}, ) raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={"code": "staff_auth_expired", "message": "Staff session expired. Please sign in again."}, ) provided = request.headers.get("x-dre-staff-key") if verify_staff_key(provided): return staff_identity(provided) raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={"code": "forbidden", "message": "Valid staff credentials required."}, ) def require_client(request: Request) -> dict: """Client session auth dependency. Returns {'client_id':..., 'client_number':...}. Raises 401 if missing/invalid/expired.""" auth = request.headers.get("authorization", "") if not auth.lower().startswith("bearer "): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail={"code": "unauthorized", "message": "Authentication required."}, ) raw_token = auth.split(" ", 1)[1].strip() if not raw_token: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail={"code": "unauthorized", "message": "Authentication required."}, ) token_hash = _sha256_hex(raw_token) now = utcnow_iso() with get_conn() as conn: row = conn.execute( "SELECT s.id, s.client_id, s.expires_at, s.revoked_at, c.client_number " "FROM sessions s JOIN clients c ON c.id = s.client_id " "WHERE s.session_hash = ?", (token_hash,), ).fetchone() if row is None: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail={"code": "unauthorized", "message": "Invalid or expired session."}, ) if row["revoked_at"] is not None: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail={"code": "unauthorized", "message": "Session revoked."}, ) if row["expires_at"] < now: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail={"code": "unauthorized", "message": "Session expired."}, ) # Touch last_seen conn.execute( "UPDATE sessions SET last_seen_at = ? WHERE id = ?", (now, row["id"]), ) conn.commit() return {"client_id": row["client_id"], "client_number": row["client_number"]} def create_magic_token(conn, client_id: str, ip: str) -> str: """Create a magic-link token row. Returns the RAW token (caller emails it, never stores it).""" raw = generate_magic_token() token_hash = _sha256_hex(raw) now = utcnow_iso() expires_at = (datetime.now(timezone.utc) + timedelta(minutes=MAGIC_TOKEN_TTL_MIN)).strftime("%Y-%m-%dT%H:%M:%SZ") conn.execute( "INSERT INTO auth_tokens (id, client_id, token_hash, expires_at, consumed_at, requested_ip, created_at) " "VALUES (?, ?, ?, ?, NULL, ?, ?)", (new_uuid(), client_id, token_hash, expires_at, ip, now), ) return raw def prune_expired_tokens(conn) -> None: """Delete auth_token rows older than 1 day (lazy sweep on verify).""" cutoff = (datetime.now(timezone.utc) - timedelta(days=1)).strftime("%Y-%m-%dT%H:%M:%SZ") conn.execute("DELETE FROM auth_tokens WHERE expires_at < ?", (cutoff,)) def verify_magic_token(conn, raw_token: str) -> str | None: """Hash token, find unexpired+unconsumed row. If found: mark consumed, return client_id. Returns None if no match.""" token_hash = _sha256_hex(raw_token) now = utcnow_iso() prune_expired_tokens(conn) row = conn.execute( "SELECT id, client_id FROM auth_tokens WHERE token_hash = ? AND consumed_at IS NULL AND expires_at > ?", (token_hash, now), ).fetchone() if row is None: return None conn.execute("UPDATE auth_tokens SET consumed_at = ? WHERE id = ?", (now, row["id"])) return row["client_id"] def create_session(conn, client_id: str) -> tuple[str, str]: """Create a session. Returns (raw_session_token, expires_at_iso).""" raw = generate_session_token() session_hash = _sha256_hex(raw) now = utcnow_iso() expires_at = (datetime.now(timezone.utc) + timedelta(days=SESSION_TTL_DAYS)).strftime("%Y-%m-%dT%H:%M:%SZ") conn.execute( "INSERT INTO sessions (id, client_id, session_hash, expires_at, revoked_at, created_at, last_seen_at) " "VALUES (?, ?, ?, ?, NULL, ?, NULL)", (new_uuid(), client_id, session_hash, expires_at, now), ) return raw, expires_at