"""JWT auth decorators + role guards + shared API helpers.""" from __future__ import annotations import functools from typing import Any, Callable from flask import g, jsonify, request from ..auth.users import AuthError from ..config import Config class ApiError(Exception): def __init__(self, message: str, status: int = 400): super().__init__(message) self.message = message self.status = status def _get_store(): from flask import current_app return current_app.extensions["user_store"] def current_user() -> dict[str, Any]: return g.user def current_org_id() -> str: return getattr(g, "org_id", None) or (g.user or {}).get("org_id") or "org-default" def assert_tenant(org_id: str | None) -> None: """Raise 403 if the object's org does not belong to the actor's tenant. super_admin is the platform operator and is exempt (can see all orgs).""" if (g.user or {}).get("role") == "super_admin": return if (org_id or "org-default") != current_org_id(): raise ApiError("permission denied", 403) def require_auth(fn: Callable) -> Callable: @functools.wraps(fn) def wrapper(*args, **kwargs): header = request.headers.get("Authorization", "") scheme, _, token = header.partition(" ") if scheme.lower() != "bearer" or not token: raise ApiError("authentication required", 401) try: payload = _get_store().decode_token(token) except AuthError as exc: raise ApiError(str(exc), 401) user = _get_store().get_user_or_none(payload.get("sub", "")) if not user or not user.get("active", True): raise ApiError("account is inactive", 401) g.user = user g.token_payload = payload # Tenant context: every request carries the actor's org id so downstream guards # can enforce per-org isolation without re-reading the user each time. g.org_id = (user or {}).get("org_id") or "org-default" return fn(*args, **kwargs) return wrapper def require_roles(*roles: str) -> Callable: def deco(fn: Callable) -> Callable: @functools.wraps(fn) def wrapper(*args, **kwargs): role = g.user.get("role") # super_admin passes any role gate allowed = {"super_admin", *roles} if role not in allowed: raise ApiError("permission denied", 403) return fn(*args, **kwargs) return wrapper return deco def api_error_handler(err: ApiError): return jsonify({"error": err.message}), err.status def register_error_handlers(app) -> None: app.register_error_handler(ApiError, api_error_handler) app.register_error_handler(ValueError, lambda e: (jsonify({"error": str(e)}), 400))