Spaces:
Running on CPU Upgrade
Running on CPU Upgrade
| from __future__ import annotations | |
| from dataclasses import dataclass, field | |
| from datetime import datetime, timezone | |
| from typing import Any | |
| from fastapi import HTTPException, Request | |
| from huggingface_hub import HfApi | |
| try: # huggingface_hub>=1.0 provides official FastAPI helpers. | |
| from huggingface_hub import parse_huggingface_oauth | |
| except Exception: # pragma: no cover - compatibility fallback for older local envs. | |
| parse_huggingface_oauth = None # type: ignore[assignment] | |
| from .security import redact | |
| REQUIRED_OAUTH_SCOPES: set[str] = { | |
| "read-repos", | |
| "write-repos", | |
| "manage-repos", | |
| "gated-repos", | |
| "inference-api", | |
| "jobs", | |
| "read-billing", | |
| } | |
| class OAuthContext: | |
| username: str | |
| token: str | |
| profile: dict[str, Any] = field(default_factory=dict) | |
| scopes: set[str] = field(default_factory=set) | |
| expires_at: datetime | None = None | |
| is_pro: bool | None = None | |
| can_pay: bool | None = None | |
| def missing_scopes(self) -> list[str]: | |
| return sorted(REQUIRED_OAUTH_SCOPES - self.scopes) | |
| def is_expired(self) -> bool: | |
| if self.expires_at is None: | |
| return False | |
| return self.expires_at <= datetime.now(timezone.utc) | |
| def _parse_scope(scope: Any) -> set[str]: | |
| if not scope: | |
| return set() | |
| if isinstance(scope, str): | |
| # HF OAuth scope strings are space-separated; be tolerant of comma lists. | |
| return {part for chunk in scope.split(",") for part in chunk.split() if part} | |
| if isinstance(scope, (list, tuple, set)): | |
| return {str(part) for part in scope if part} | |
| return {str(scope)} | |
| def _normalize_expires_at(value: Any) -> datetime | None: | |
| if value is None: | |
| return None | |
| if isinstance(value, datetime): | |
| if value.tzinfo is None: | |
| return value.replace(tzinfo=timezone.utc) | |
| return value.astimezone(timezone.utc) | |
| try: | |
| return datetime.fromtimestamp(float(value), tz=timezone.utc) | |
| except Exception: | |
| return None | |
| def _ctx_from_official_parser(request: Request) -> OAuthContext | None: | |
| if parse_huggingface_oauth is None: | |
| return None | |
| try: | |
| info = parse_huggingface_oauth(request) # type: ignore[misc] | |
| except AssertionError: | |
| # SessionMiddleware is not present in local/custom-only fallback mode. | |
| return None | |
| if info is None: | |
| return None | |
| user_info = getattr(info, "user_info", None) | |
| username = getattr(user_info, "preferred_username", None) or getattr(user_info, "name", None) | |
| token = getattr(info, "access_token", None) | |
| if not username or not token: | |
| return None | |
| profile = { | |
| "name": getattr(user_info, "name", None), | |
| "preferred_username": getattr(user_info, "preferred_username", None), | |
| "picture": getattr(user_info, "picture", None), | |
| "email": getattr(user_info, "email", None), | |
| "is_pro": getattr(user_info, "is_pro", None), | |
| "can_pay": getattr(user_info, "can_pay", None), | |
| } | |
| return OAuthContext( | |
| username=str(username), | |
| token=str(token), | |
| profile={k: v for k, v in profile.items() if v is not None}, | |
| scopes=_parse_scope(getattr(info, "scope", None)), | |
| expires_at=_normalize_expires_at(getattr(info, "access_token_expires_at", None)), | |
| is_pro=getattr(user_info, "is_pro", None), | |
| can_pay=getattr(user_info, "can_pay", None), | |
| ) | |
| def _ctx_from_raw_session(request: Request) -> OAuthContext | None: | |
| try: | |
| oauth_info = request.session.get("oauth_info") # type: ignore[attr-defined] | |
| except Exception: | |
| oauth_info = None | |
| if not oauth_info: | |
| return None | |
| userinfo = oauth_info.get("userinfo") or {} | |
| username = userinfo.get("preferred_username") or userinfo.get("username") or userinfo.get("name") | |
| token = oauth_info.get("access_token") | |
| if not username or not token: | |
| return None | |
| profile = { | |
| "name": userinfo.get("name"), | |
| "preferred_username": userinfo.get("preferred_username") or userinfo.get("username"), | |
| "picture": userinfo.get("picture"), | |
| "email": userinfo.get("email"), | |
| "is_pro": userinfo.get("isPro") or userinfo.get("is_pro"), | |
| "can_pay": userinfo.get("canPay") or userinfo.get("can_pay"), | |
| } | |
| return OAuthContext( | |
| username=str(username), | |
| token=str(token), | |
| profile={k: v for k, v in profile.items() if v is not None}, | |
| scopes=_parse_scope(oauth_info.get("scope")), | |
| expires_at=_normalize_expires_at(oauth_info.get("expires_at")), | |
| is_pro=profile.get("is_pro"), | |
| can_pay=profile.get("can_pay"), | |
| ) | |
| def extract_oauth_context(request: Request) -> OAuthContext: | |
| """Extract and validate the signed-in HF user from the Gradio/FastAPI OAuth session. | |
| Uses the official `huggingface_hub.parse_huggingface_oauth` helper first, then falls | |
| back to the raw Gradio session shape for compatibility. The token is kept server-side | |
| only and must never be returned by API responses. | |
| """ | |
| ctx = _ctx_from_official_parser(request) or _ctx_from_raw_session(request) | |
| if ctx is None: | |
| raise HTTPException(status_code=401, detail="Please sign in with Hugging Face first.") | |
| if ctx.is_expired: | |
| raise HTTPException(status_code=401, detail="Your Hugging Face OAuth session expired. Please sign in again.") | |
| return ctx | |
| def public_oauth_context(ctx: OAuthContext) -> dict[str, Any]: | |
| return { | |
| "username": ctx.username, | |
| "profile": { | |
| "name": ctx.profile.get("name"), | |
| "preferred_username": ctx.profile.get("preferred_username") or ctx.username, | |
| "picture": ctx.profile.get("picture"), | |
| "is_pro": ctx.is_pro, | |
| "can_pay": ctx.can_pay, | |
| }, | |
| "scopes": sorted(ctx.scopes), | |
| "missing_scopes": ctx.missing_scopes, | |
| "expires_at": ctx.expires_at.isoformat() if ctx.expires_at else None, | |
| "authenticated": True, | |
| } | |
| def oauth_warning_messages(ctx: OAuthContext) -> list[str]: | |
| warnings: list[str] = [] | |
| if ctx.missing_scopes: | |
| warnings.append("Missing OAuth scopes: " + ", ".join(ctx.missing_scopes)) | |
| if ctx.can_pay is False: | |
| warnings.append("No billing/payment method is visible through OAuth; fixed GPU hardware may require manual action.") | |
| return warnings | |
| def verify_token_identity(ctx: OAuthContext) -> dict[str, Any]: | |
| """Best-effort diagnostics endpoint helper. Never returns the raw token.""" | |
| try: | |
| info = HfApi(token=ctx.token).whoami() | |
| name = info.get("name") or info.get("fullname") or info.get("preferred_username") | |
| return { | |
| "ok": True, | |
| "oauth_username": ctx.username, | |
| "whoami_name": name, | |
| "matches_oauth_user": name == ctx.username if name else None, | |
| "can_pay": ctx.can_pay, | |
| "is_pro": ctx.is_pro, | |
| "missing_scopes": ctx.missing_scopes, | |
| } | |
| except Exception as exc: # noqa: BLE001 | |
| return { | |
| "ok": False, | |
| "oauth_username": ctx.username, | |
| "error": redact(str(exc)), | |
| "missing_scopes": ctx.missing_scopes, | |
| } | |