RL-Explorer / app /auth.py
AdithyaSK's picture
AdithyaSK HF Staff
Deploy HF RL Explorer
da5cba1 verified
Raw History Blame Contribute Delete
11.3 kB
"""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