Spaces:
Running on Zero
Running on Zero
| import hmac | |
| import json | |
| import logging | |
| from typing import Dict, List, Optional, Any | |
| from urllib.parse import urlparse | |
| import requests | |
| from fastapi import Header, HTTPException | |
| from starlette.responses import JSONResponse | |
| from app.core.config import ( | |
| SUPABASE_URL, SUPABASE_ANON_KEY, SUPABASE_SERVICE_ROLE_KEY, SUPABASE_ENFORCE_AUTH, | |
| SUPABASE_TIMEOUT_SECONDS, BARRIER_API_TOKENS, BARRIER_SYNTHETIC_USER_ID, | |
| BARRIER_ROLES, REDIRECT_ALLOWED_HTTP_HOSTS, REDIRECT_ALLOWED_DEEPLINK_SCHEMES, | |
| MOBILE_REDIRECT_SCHEME_PATTERN, PASSWORD_RESET_REDIRECT_URL, INTERNAL_LOGIN_EMAIL_SUFFIX, | |
| PHONE_PATTERN, USERNAME_PATTERN, NATIONAL_ID_PATTERN, | |
| supabase_configured, is_staff_role, is_barrier_role, can_view_global_records, | |
| is_developer_role, LOGGER, | |
| ) | |
| LOGGER = logging.getLogger("ain_el_aql.security") | |
| def _normalize_spaces(value: str) -> str: | |
| return " ".join((value or "").split()) | |
| def _extract_bearer_token(authorization: Optional[str]) -> Optional[str]: | |
| if not authorization: | |
| return None | |
| parts = authorization.strip().split(" ", 1) | |
| if len(parts) != 2 or parts[0].lower() != "bearer": | |
| return None | |
| token = parts[1].strip() | |
| return token if token else None | |
| def _looks_like_jwt(token: str) -> bool: | |
| parts = token.split(".") | |
| if len(parts) != 3: | |
| return False | |
| return bool(parts[0] and parts[1]) | |
| def _extract_response_error_message(response: requests.Response) -> str: | |
| try: | |
| payload = response.json() | |
| except ValueError: | |
| payload = None | |
| if isinstance(payload, dict): | |
| for key in ["error_description", "message", "error", "msg"]: | |
| value = payload.get(key) | |
| if value is not None: | |
| text = str(value).strip() | |
| if text: | |
| return text | |
| return response.text[:250].strip() | |
| def _supabase_headers(*, api_key: str, bearer: str, content_type: Optional[str] = None, prefer: Optional[str] = None) -> Dict[str, str]: | |
| headers: Dict[str, str] = { | |
| "apikey": api_key, | |
| "Authorization": f"Bearer {bearer}", | |
| } | |
| if content_type: | |
| headers["Content-Type"] = content_type | |
| if prefer: | |
| headers["Prefer"] = prefer | |
| return headers | |
| def verify_barrier_token(token: str) -> Optional[Dict[str, Any]]: | |
| if not token or not BARRIER_API_TOKENS: | |
| return None | |
| for configured_token in BARRIER_API_TOKENS: | |
| if hmac.compare_digest(token, configured_token): | |
| return { | |
| "id": BARRIER_SYNTHETIC_USER_ID, | |
| "email": None, | |
| "role": "barrier", | |
| "auth_type": "barrier_token", | |
| } | |
| return None | |
| def verify_supabase_user_token(token: str) -> Dict[str, Any]: | |
| if not supabase_configured(): | |
| raise HTTPException(status_code=503, detail="Supabase auth is not configured on backend.") | |
| if not _looks_like_jwt(token): | |
| raise HTTPException(status_code=401, detail="Invalid Supabase access token format. Send session.access_token as Bearer token.") | |
| url = f"{SUPABASE_URL}/auth/v1/user" | |
| try: | |
| response = requests.get( | |
| url, | |
| headers={"apikey": SUPABASE_ANON_KEY, "Authorization": f"Bearer {token}"}, | |
| timeout=SUPABASE_TIMEOUT_SECONDS, | |
| ) | |
| except requests.RequestException as exc: | |
| raise HTTPException(status_code=503, detail=f"Supabase auth is unreachable: {exc}") from exc | |
| if response.status_code != 200: | |
| supabase_error = _extract_response_error_message(response) | |
| err_lc = supabase_error.lower() | |
| if "expired" in err_lc: | |
| detail = "Supabase access token expired. Refresh session and retry." | |
| elif any(hint in err_lc for hint in ["jwt", "signature", "malformed", "invalid"]): | |
| detail = "Invalid Supabase access token." | |
| else: | |
| detail = "Invalid or expired Supabase token." | |
| raise HTTPException(status_code=401, detail=detail) | |
| payload = response.json() | |
| if not isinstance(payload, dict): | |
| raise HTTPException(status_code=401, detail="Invalid Supabase auth response payload.") | |
| user_id = payload.get("id") | |
| if not user_id: | |
| raise HTTPException(status_code=401, detail="Supabase token did not return a user id.") | |
| # Fetch role from Supabase profiles table | |
| role = "user" | |
| try: | |
| from app.services.supabase_client import supabase_get_profile_by_user_id | |
| profile = supabase_get_profile_by_user_id(user_id) | |
| if isinstance(profile, dict) and profile.get("role"): | |
| role = str(profile.get("role")).strip().lower() | |
| except Exception: | |
| pass | |
| return { | |
| "id": user_id, | |
| "email": payload.get("email"), | |
| "role": role, | |
| "auth_type": "supabase", | |
| "user_metadata": payload.get("user_metadata") if isinstance(payload.get("user_metadata"), dict) else {}, | |
| "app_metadata": payload.get("app_metadata") if isinstance(payload.get("app_metadata"), dict) else {}, | |
| } | |
| def require_authenticated_user(authorization: Optional[str]) -> Optional[Dict[str, Any]]: | |
| token = _extract_bearer_token(authorization) | |
| if SUPABASE_ENFORCE_AUTH and not token: | |
| raise HTTPException(status_code=401, detail="Missing Bearer token.") | |
| if not token: | |
| return None | |
| barrier_user = verify_barrier_token(token) | |
| if barrier_user is not None: | |
| return barrier_user | |
| return verify_supabase_user_token(token) | |
| def require_developer_user(authorization: Optional[str]): | |
| request_user = require_authenticated_user(authorization) | |
| if request_user is None: | |
| raise HTTPException(status_code=401, detail="Authentication is required.") | |
| requester_role = (request_user.get("role") or "").strip().lower() | |
| if not is_developer_role(requester_role): | |
| raise HTTPException(status_code=403, detail="Only developers can access this endpoint.") | |
| return request_user, requester_role | |
| def error_response(status_code: int, message: str, code: str) -> JSONResponse: | |
| return JSONResponse( | |
| status_code=status_code, | |
| content={"status": "error", "message": message, "code": code}, | |
| ) | |
| def is_valid_redirect_url(url: str, *, allowed_http_hosts: Optional[List[str]] = None) -> bool: | |
| if not url: | |
| return False | |
| parsed = urlparse(url) | |
| scheme = (parsed.scheme or "").strip() | |
| if not scheme or not MOBILE_REDIRECT_SCHEME_PATTERN.fullmatch(scheme): | |
| return False | |
| normalized_scheme = scheme.lower() | |
| if normalized_scheme in {"http", "https"}: | |
| hostname = (parsed.hostname or "").strip().lower() | |
| if not hostname: | |
| return False | |
| effective_allowed_hosts = set(REDIRECT_ALLOWED_HTTP_HOSTS) | |
| if allowed_http_hosts: | |
| for host in allowed_http_hosts: | |
| normalized_host = (host or "").strip().lower() | |
| if normalized_host: | |
| effective_allowed_hosts.add(normalized_host) | |
| if not effective_allowed_hosts: | |
| return False | |
| return hostname in effective_allowed_hosts | |
| if normalized_scheme not in REDIRECT_ALLOWED_DEEPLINK_SCHEMES: | |
| return False | |
| return bool(parsed.netloc or parsed.path) | |