"""Who the user is: Sign in with Hugging Face (OAuth), or an access token they paste. Either way the session is one encrypted, HttpOnly cookie. The token inside is read per request to start the user's own sandbox and call Inference Providers as them. It is never written to disk or into a trace, and the cookie is encrypted rather than only signed, so the token cannot be read back out of it. Locally there is no OAuth app: your own token (HF_TOKEN or `hf auth login`) is the user, unless you paste a different one. """ from __future__ import annotations import base64 import hashlib import json import os import secrets import time from urllib.parse import urlencode import httpx from cryptography.fernet import Fernet, InvalidToken from fastapi import APIRouter, HTTPException, Request from fastapi.responses import JSONResponse, RedirectResponse from pydantic import BaseModel, Field from . import config from .http import Limiter, client_ip router = APIRouter() COOKIE = "rlx_session" _box = Fernet(base64.urlsafe_b64encode(hashlib.sha256(f"rl-explorer:{config.SESSION_SECRET}".encode()).digest())) STATE_COOKIE = "rlx_oauth" # the sign-in's state, encrypted with its time: no server memory, any replica, any restart STATE_TTL = 600 SIGNIN_LIMIT = Limiter(30, 600) # sign-ins started per address TOKEN_LIMIT = Limiter(10, 600) # pasted tokens checked per address: each is a call to the Hub on its behalf _local: dict = {} TOKENS_URL = "https://huggingface.co/settings/tokens" def _redirect_uri(request: Request) -> str: host = config.SPACE_HOST or request.url.netloc scheme = "https" if config.SPACE_HOST else request.url.scheme return f"{scheme}://{host}/login/callback" def _local_user() -> dict | None: if _local: return _local from huggingface_hub import get_token, whoami token = get_token() if not token: return None try: me = whoami(token=token) except Exception: return None _local.update(token=token, name=me["name"], avatar=me.get("avatarUrl"), local=True, via="local", orgs=_orgs(me)) return _local def _orgs(info: dict) -> list[str]: """The organizations the Hub says this account belongs to (whoami-v2 or OAuth userinfo), private memberships too.""" return sorted({o.get("name") or o.get("preferred_username") for o in (info.get("orgs") or []) if isinstance(o, dict)} - {None})[:100] def _cookie_user(request: Request) -> dict | None: raw = request.cookies.get(COOKIE) if not raw: return None try: data = json.loads(_box.decrypt(raw.encode(), ttl=config.SESSION_DAYS * 86400)) except (InvalidToken, ValueError): return None return data if data.get("exp", 0) > time.time() else None TRUST_NETWORK = os.environ.get("RLX_TRUST_NETWORK") == "1" def _is_loopback(request: Request) -> bool: import ipaddress # a browser on this machine talks to us directly; forwarding headers mean a proxy is in between (or someone # spoofing one, which uvicorn's --proxy-headers would otherwise believe), so the peer isn't provably local if any(h in request.headers for h in ("x-forwarded-for", "x-real-ip", "forwarded")): return False try: return ipaddress.ip_address((request.client.host if request.client else "") or "").is_loopback except ValueError: return False def current_user(request: Request) -> dict | None: """Locally, the machine's own token signs in only requests from this machine (a browser on it, or `docker run -p 127.0.0.1:...` with RLX_TRUST_NETWORK=1): anyone else on the network has to sign in.""" u = _cookie_user(request) if u or not config.LOCAL_MODE: return u return _local_user() if (_is_loopback(request) or TRUST_NETWORK) else None _billing_cache: dict[str, tuple[float, dict]] = {} def billing(u: dict) -> dict: """{"can_pay": bool | None, "mode": "prepaid" | ...} from the Hub for this account, cached for 5 minutes.""" token = u.get("token") if not token: return {"can_pay": None} key = u["name"] hit = _billing_cache.get(key) if hit and time.time() - hit[0] < 300: return hit[1] try: r = httpx.get(f"{config.OPENID_PROVIDER_URL}/api/whoami-v2", headers={"Authorization": f"Bearer {token}"}, timeout=10) me = r.json() if r.status_code == 200 else {} except (httpx.HTTPError, ValueError): me = {} out = {"can_pay": me.get("canPay"), "mode": me.get("billingMode")} if len(_billing_cache) > 5000: _billing_cache.clear() _billing_cache[key] = (time.time(), out) return out def require_user(request: Request) -> dict: u = current_user(request) if not u: raise HTTPException(401, "Sign in with Hugging Face to run rollouts.") return u def public(u: dict | None) -> dict | None: if not u: return None return {"name": u["name"], "avatar": u.get("avatar"), "local": bool(u.get("local")), "via": u.get("via", "oauth")} def _issue(resp, session: dict): # SameSite=None so it also works inside the huggingface.co iframe. Browsers accept Secure on http://localhost. resp.set_cookie(COOKIE, _box.encrypt(json.dumps(session).encode()).decode(), httponly=True, secure=True, samesite="none", max_age=config.SESSION_DAYS * 86400, path="/") return resp # ── OAuth ──────────────────────────────────────────────────────────────────── @router.get("/login") def login(request: Request): if config.LOCAL_MODE: return RedirectResponse("/") if not SIGNIN_LIMIT.hit(client_ip(request)): raise HTTPException(429, "Too many sign-ins from here. Try again in a few minutes.") state = secrets.token_urlsafe(24) q = urlencode({"client_id": config.OAUTH_CLIENT_ID, "redirect_uri": _redirect_uri(request), "response_type": "code", "scope": " ".join(config.OAUTH_SCOPES), "state": state}) resp = RedirectResponse(f"{config.OPENID_PROVIDER_URL}/oauth/authorize?{q}") resp.set_cookie(STATE_COOKIE, _box.encrypt(state.encode()).decode(), httponly=True, secure=True, samesite="none", max_age=STATE_TTL, path="/login") return resp def _state_ok(request: Request, state: str) -> bool: """The callback's state is the one this browser's sign-in started with, and that sign-in is under 10 minutes old.""" try: expected = _box.decrypt(request.cookies.get(STATE_COOKIE, "").encode(), ttl=STATE_TTL).decode() except (InvalidToken, ValueError): return False return bool(state) and secrets.compare_digest(state.encode(), expected.encode()) @router.get("/login/callback") def callback(request: Request, code: str = "", state: str = "", error: str = ""): if error or not code: # the visitor said no on huggingface.co resp = RedirectResponse("/") resp.delete_cookie(STATE_COOKIE, path="/login", samesite="none", secure=True) return resp if not _state_ok(request, state): raise HTTPException(400, "Sign-in expired or was tampered with. Try again.") basic = base64.b64encode(f"{config.OAUTH_CLIENT_ID}:{config.OAUTH_CLIENT_SECRET}".encode()).decode() r = httpx.post(f"{config.OPENID_PROVIDER_URL}/oauth/token", timeout=20, headers={"Authorization": f"Basic {basic}"}, data={"grant_type": "authorization_code", "code": code, "redirect_uri": _redirect_uri(request), "client_id": config.OAUTH_CLIENT_ID}) if r.status_code != 200: raise HTTPException(400, f"Could not complete sign-in ({r.status_code}).") tok = r.json() granted = set((tok.get("scope") or "").split()) info = httpx.get(f"{config.OPENID_PROVIDER_URL}/oauth/userinfo", timeout=20, headers={"Authorization": f"Bearer {tok['access_token']}"}).json() resp = RedirectResponse("/") resp.delete_cookie(STATE_COOKIE, path="/login", samesite="none", secure=True) # one use return _issue(resp, { "token": tok["access_token"], "name": info.get("preferred_username") or info.get("name"), "avatar": info.get("picture"), "exp": time.time() + min(tok.get("expires_in") or 28800, config.SESSION_DAYS * 86400), "via": "oauth", "orgs": _orgs(info), "missing_scopes": sorted({"inference-api", "jobs"} - granted)}) # ── access token ───────────────────────────────────────────────────────────── class TokenLogin(BaseModel): token: str = Field(max_length=300) def _token_gaps(access: dict) -> list[str]: """What a token can't do that a rollout needs. Write tokens can do everything; read tokens can't start sandboxes; fine-grained tokens are checked against their global permissions (a warning, not a refusal).""" role = access.get("role") if role == "write": return [] if role == "read": return ["jobs"] perms = " ".join((access.get("fineGrained") or {}).get("global") or []) return [n for n, key in (("inference-api", "inference"), ("jobs", "job")) if key not in perms] @router.post("/api/login/token") def token_login(body: TokenLogin, request: Request): if not TOKEN_LIMIT.hit(client_ip(request)): raise HTTPException(429, "Too many tokens tried from here. Wait a few minutes.") token = body.token.strip() if not token.startswith("hf_") or len(token) < 20 or any(c.isspace() for c in token): raise HTTPException(400, "That doesn't look like a Hugging Face token. They start with hf_.") r = httpx.get(f"{config.OPENID_PROVIDER_URL}/api/whoami-v2", headers={"Authorization": f"Bearer {token}"}, timeout=20) if r.status_code == 401: raise HTTPException(400, "Hugging Face didn't accept this token. It may have been revoked or mistyped.") if r.status_code != 200: raise HTTPException(502, f"Couldn't check the token with Hugging Face ({r.status_code}). Try again.") me = r.json() if me.get("type") != "user": raise HTTPException(400, "Use a token that belongs to your user account, not an organization token.") gaps = _token_gaps(((me.get("auth") or {}).get("accessToken")) or {}) if gaps == ["jobs"] and (me["auth"]["accessToken"].get("role") == "read"): raise HTTPException(400, "This is a read token, and rollouts need to start an HF Sandbox. Use a write token, or a " "fine-grained one with the Inference Providers and Jobs permissions.") resp = JSONResponse({"user": {"name": me["name"], "avatar": me.get("avatarUrl"), "local": False, "via": "token"}, "missing_scopes": gaps}) return _issue(resp, {"token": token, "name": me["name"], "avatar": me.get("avatarUrl"), "via": "token", "orgs": _orgs(me), "exp": time.time() + config.SESSION_DAYS * 86400, "missing_scopes": gaps}) @router.post("/api/logout") def logout_api(): resp = JSONResponse({"ok": True}) resp.delete_cookie(COOKIE, samesite="none", secure=True, path="/") return resp