Spaces:
Running
Running
Download app/auth.py from FineEnvs/RL-Explorer: direct link, hf CLI and curl.
- Browser
- Download file 11.3 kB
-
https://huggingface.co/spaces/FineEnvs/RL-Explorer/resolve/main/app/auth.py
- Command line
-
hf download hf://spaces/FineEnvs/RL-Explorer/app/auth.py
-
curl -L -o auth.py https://huggingface.co/spaces/FineEnvs/RL-Explorer/resolve/main/app/auth.py
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 ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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()) | |
| 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] | |
| 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}) | |
| def logout_api(): | |
| resp = JSONResponse({"ok": True}) | |
| resp.delete_cookie(COOKIE, samesite="none", secure=True, path="/") | |
| return resp | |