"""Same-origin API and server-side OAuth for the Annotation Studio Docker Space.""" import base64 import hashlib import json import os import re import secrets import tempfile import threading import time from pathlib import Path from typing import Literal from urllib.parse import urlencode import requests from fastapi import FastAPI, HTTPException, Request from fastapi.responses import FileResponse, JSONResponse, RedirectResponse, Response from fastapi.staticfiles import StaticFiles from huggingface_hub import HfApi, hf_hub_download from model_catalog import DEFAULT_MODEL, public_models from prompt_config import public_catalog, validate_prompts from pydantic import BaseModel, Field from reporting import annotation_url, usage_costs from studio import COPY, FEATURES, launch, prepare, read_cameras app = FastAPI(title="LeRobot Annotation Studio", docs_url=None, redoc_url=None, openapi_url=None) BASE = Path(__file__).parent HOST = os.environ.get("SPACE_HOST") ORIGIN = f"https://{HOST}" if HOST else os.environ.get("STUDIO_LOCAL_ORIGIN", "http://127.0.0.1:7860") COOKIE = "annotation_session" SESSIONS = {} QUOTES = {} LOCK = threading.RLock() ESTIMATES = threading.BoundedSemaphore(4) LAUNCH_LOCK = threading.Lock() class DatasetInput(BaseModel): dataset: str = Field(min_length=3, max_length=200) class EstimateInput(DatasetInput): model: str = Field(default=DEFAULT_MODEL, min_length=1, max_length=80) start: int = Field(default=0, ge=0) count: int = Field(default=100, ge=0) features: list[str] = Field(min_length=1, max_length=7) camera: str | None = Field(default=None, max_length=200) hz: float = 0.1 hours: int = 4 mode: Literal["Annotated dataset copy"] = COPY prompts: dict[str, str] = Field(default_factory=dict, max_length=7) # 0 = every subtask. human_video_segments: int | None = Field(default=None, ge=0, le=8) human_video_parallel: int = Field(default=8, ge=1, le=32) class PromptInput(BaseModel): prompts: dict[str, str] = Field(default_factory=dict, max_length=7) class LaunchInput(BaseModel): quote_id: str = Field(min_length=16, max_length=100) output: str = Field(min_length=3, max_length=200) budget: float = Field(ge=1, le=5000) confirmed: bool visibility: Literal["private", "public"] = "private" def clean_expired(): now = time.time() with LOCK: for key in list(SESSIONS): if SESSIONS[key]["expires"] < now: del SESSIONS[key] for key in list(QUOTES): if QUOTES[key]["created"] < now - 1800: del QUOTES[key] def session(request, required=False): clean_expired() value = SESSIONS.get(request.cookies.get(COOKIE, "")) if required and (not value or not value.get("token")): raise HTTPException(401, "Sign in with Hugging Face to continue.") return value def create_session(): clean_expired() with LOCK: if len(SESSIONS) > 4096: raise HTTPException(503, "The workspace is busy. Please try again shortly.") sid = secrets.token_urlsafe(32) SESSIONS[sid] = {"csrf": secrets.token_urlsafe(32), "expires": time.time() + 86400} return sid, SESSIONS[sid] def cookie(response, sid): response.set_cookie(COOKIE, sid, httponly=True, secure=bool(HOST), samesite="lax", max_age=86400) return response def check_origin(request): allowed = {ORIGIN} if not HOST: allowed.add("http://localhost:7860") if request.headers.get("origin") not in allowed: raise HTTPException(403, "Open the workspace in its own tab and try again.") def csrf(request): check_origin(request) current = session(request, required=True) if not secrets.compare_digest(request.headers.get("x-csrf-token", ""), current["csrf"]): raise HTTPException(403, "Your session changed. Refresh the workspace.") return current def safe_error(exc): if isinstance(exc, ValueError): return str(exc) code = getattr(getattr(exc, "response", None), "status_code", None) if code in (401, 403): return "Hugging Face denied access. Check your permissions and available credits." if code == 404: return "Dataset not found or inaccessible. Check the ID; sign in for private or gated datasets." if code == 409: return "This output dataset already exists. Choose a new name." return "Hugging Face could not complete this request. Please retry or check your Jobs page." @app.middleware("http") async def security_headers(request, call_next): response = await call_next(request) response.headers["X-Content-Type-Options"] = "nosniff" response.headers["Referrer-Policy"] = "strict-origin-when-cross-origin" if request.url.path.startswith(("/api", "/auth")): response.headers["Cache-Control"] = "no-store" return response @app.get("/api/session") def get_session(request: Request): current = session(request) sid = None if current is None: sid, current = create_session() response = JSONResponse( { "user": current.get("user"), "csrf": current["csrf"], "auth_available": bool(os.environ.get("OAUTH_CLIENT_ID") and HOST), "features": list(FEATURES), } ) return cookie(response, sid) if sid else response @app.get("/auth/login") def login(request: Request): client_id = os.environ.get("OAUTH_CLIENT_ID") if not client_id or not HOST: return RedirectResponse("/?error=local-preview") sid, current = create_session() current["oauth_state"] = secrets.token_urlsafe(32) current["verifier"] = secrets.token_urlsafe(64) current["auth_started"] = time.time() # Authorization Code + PKCE. Tokens never pass through frontend JavaScript. challenge = ( base64.urlsafe_b64encode(hashlib.sha256(current["verifier"].encode()).digest()).rstrip(b"=").decode() ) params = { "client_id": client_id, "redirect_uri": ORIGIN + "/auth/callback", "response_type": "code", "scope": os.environ.get( "OAUTH_SCOPES", "openid profile read-repos contribute-repos inference-api jobs" ), "state": current["oauth_state"], "code_challenge": challenge, "code_challenge_method": "S256", } return cookie(RedirectResponse("https://huggingface.co/oauth/authorize?" + urlencode(params)), sid) @app.get("/auth/callback") def callback(request: Request, code: str = "", state: str = "", error: str = ""): current = session(request) if ( error or not code or not current or time.time() - current.get("auth_started", 0) > 600 or not secrets.compare_digest(state, current.get("oauth_state", "invalid")) ): return RedirectResponse("/?error=sign-in") verifier = current.pop("verifier", "") current.pop("oauth_state", None) try: result = requests.post( "https://huggingface.co/oauth/token", data={ "grant_type": "authorization_code", "code": code, "redirect_uri": ORIGIN + "/auth/callback", "code_verifier": verifier, }, auth=(os.environ["OAUTH_CLIENT_ID"], os.environ["OAUTH_CLIENT_SECRET"]), timeout=30, ) result.raise_for_status() data = result.json() token = data["access_token"] user = HfApi(token=token).whoami() except Exception: return RedirectResponse("/?error=sign-in") with LOCK: SESSIONS.pop(request.cookies.get(COOKIE, ""), None) sid, updated = create_session() updated.update( token=token, user={"name": user["name"], "fullname": user.get("fullname") or user["name"]}, expires=time.time() + min(int(data.get("expires_in", 86400)), 86400), ) return cookie(RedirectResponse("/"), sid) @app.post("/auth/logout") def logout(request: Request): csrf(request) with LOCK: SESSIONS.pop(request.cookies.get(COOKIE, ""), None) response = JSONResponse({"ok": True}) response.delete_cookie(COOKIE, secure=bool(HOST), samesite="lax") return response @app.post("/api/cameras") def cameras(options: DatasetInput, request: Request): check_origin(request) current = session(request) or {} if not ESTIMATES.acquire(blocking=False): raise HTTPException(429, "Other datasets are loading. Try again in a moment.") try: return read_cameras(options.dataset, current.get("token", False)) except Exception as exc: raise HTTPException(400, safe_error(exc)) from None finally: ESTIMATES.release() @app.post("/api/estimate") def calculate(options: EstimateInput, request: Request): check_origin(request) current = session(request) or {} if not ESTIMATES.acquire(blocking=False): raise HTTPException(429, "Other estimates are loading. Try again in a moment.") try: plan = prepare( options.dataset, options.start, options.count, options.features, options.camera, options.hz, options.hours, options.mode, current.get("token", False), prompts=options.prompts, model_key=options.model, human_video_segments=options.human_video_segments, human_video_parallel=options.human_video_parallel, ) except Exception as exc: raise HTTPException(400, safe_error(exc)) from None finally: ESTIMATES.release() quote_id = secrets.token_urlsafe(24) plan["quote_owner"] = (current.get("user") or {}).get("name") s = plan["selection"] # A quote needs a reproducible selection digest, not every episode's metadata. # Keep large-dataset estimates from occupying hundreds of MB for 30 minutes. plan["selection"] = {k: s[k] for k in ("dataset", "revision", "start", "count")} plan["selection"]["episode_ids_sha256"] = hashlib.sha256( json.dumps([row["episode_index"] for row in s["episodes"]]).encode() ).hexdigest() with LOCK: clean_expired() if len(QUOTES) >= 4096: raise HTTPException(503, "The estimate service is busy. Try again shortly.") QUOTES[quote_id] = plan return { "quote_id": quote_id, "costs": plan["costs"], "model_profile": plan["model_profile"], "camera": plan["camera"], "cameras": s["cameras"], "episodes": len(s["episodes"]), "total_episodes": s["total_episodes"], "seconds": s["seconds"], "fps": s["info"]["fps"], "dataset": s["dataset"], "revision": s["revision"], "private": s["private"], "first_episode": s["episodes"][0]["episode_index"], "last_episode": s["episodes"][-1]["episode_index"], "expires_at": plan["created"] + 1800, "lengths": [r["length"] / s["info"]["fps"] for r in s["episodes"][:100]], "visualizer_url": f"https://lerobot-visualize-dataset.hf.space/{s['dataset']}/{s['episodes'][0]['episode_index']}", } @app.get("/api/prompts") def prompts(): return public_catalog() @app.post("/api/prompts/validate") def check_prompts(options: PromptInput, request: Request): check_origin(request) try: return validate_prompts(options.prompts) except ValueError as exc: raise HTTPException(400, str(exc)) from None @app.post("/api/jobs") def submit(options: LaunchInput, request: Request): current = csrf(request) with LAUNCH_LOCK: plan = QUOTES.get(options.quote_id) if not plan or plan.get("quote_owner") not in (None, current["user"]["name"]): raise HTTPException(400, "Estimate expired. Calculate a fresh estimate.") if current["expires"] < time.time() + plan["costs"]["hours"] * 3600 + 900: raise HTTPException(401, "Sign in again so your login lasts for the entire job.") if plan.get("submitted"): return plan["submitted"] try: job = launch( plan, options.output, options.budget, options.confirmed, current["token"], visibility=options.visibility, ) except Exception as exc: raise HTTPException(400, safe_error(exc)) from None plan["submitted"] = job return job def output_repo(job, owner): value = (job.labels or {}).get("annotation_output", "") return value if not value or "/" in value else f"{owner}/{value}" @app.get("/api/jobs") def list_jobs(request: Request): current = session(request, required=True) try: api = HfApi(token=current["token"]) jobs = list( api.list_jobs(namespace=current["user"]["name"], labels={"app": "lerobot-annotation-studio"}) ) return [ { "id": j.id, "name": getattr(j, "name", None) or (j.labels or {}).get("name"), "stage": j.status.stage, "created_at": str(j.created_at), "owner": current["user"]["name"], "output": output_repo(j, current["user"]["name"]), } for j in jobs[:40] ] except Exception as exc: raise HTTPException(400, safe_error(exc)) from None @app.post("/api/jobs/{job_id}/cancel") def cancel(job_id: str, request: Request): current = csrf(request) api = HfApi(token=current["token"]) try: job = api.inspect_job(job_id=job_id, namespace=current["user"]["name"]) if (job.labels or {}).get("app") != "lerobot-annotation-studio": raise ValueError("This job was not started by Annotation Studio.") api.cancel_job(job_id=job_id, namespace=current["user"]["name"]) except Exception as exc: raise HTTPException(400, safe_error(exc)) from None return {"ok": True} def studio_output(job_id, current): """The job and its output dataset, only for Annotation Studio jobs in the user's namespace.""" api = HfApi(token=current["token"]) job = api.inspect_job(job_id=job_id, namespace=current["user"]["name"]) output = output_repo(job, current["user"]["name"]) if (job.labels or {}).get("app") != "lerobot-annotation-studio" or not output.startswith( current["user"]["name"] + "/" ): raise ValueError("This job has no Annotation Studio result.") return job, output def read_output_file(output, path, token): with tempfile.TemporaryDirectory() as temp: local = hf_hub_download(output, path, repo_type="dataset", token=token, cache_dir=temp) return Path(local).read_bytes() HUMAN_VIDEO_FILE = re.compile( r"annotation_studio/human_videos/episode_\d{6}/segment_\d{3}(_robot_frame\.png|_first_frame\.png|\.mp4)" ) @app.get("/api/jobs/{job_id}/progress") def progress(job_id: str, request: Request): current = session(request, required=True) try: job, output = studio_output(job_id, current) report = json.loads(read_output_file(output, "annotation_studio/run.json", current["token"])) episodes = report.get("episodes", []) videos = {} for episode in episodes: for status, count in (episode.get("human_videos") or {}).items(): videos[status] = videos.get(status, 0) + count return { "status": "interrupted" if job.status.stage in {"COMPLETED", "ERROR", "CANCELED", "DELETED"} and report["status"] == "running" else report["status"], "expected": report["expected_episodes"], "completed": sum(e["status"] == "completed" for e in episodes), "failed": sum(e["status"] == "failed" for e in episodes), "stopped": sum(e["status"] in {"time_stopped", "budget_stopped"} for e in episodes), "pending": max(0, report["expected_episodes"] - len(episodes)), "usage": report.get("usage", {}), "cost": usage_costs(job, report), "annotation_url": annotation_url(report), "output_mode": report.get("manifest", {}).get("mode"), "model": report.get("manifest", {}).get("model"), "human_videos": videos if (report.get("manifest", {}).get("human_video")) else None, } except Exception as exc: raise HTTPException(400, "No saved progress yet. The job may still be starting.") from exc @app.get("/api/jobs/{job_id}/human-videos") def human_videos(job_id: str, request: Request): current = session(request, required=True) try: _, output = studio_output(job_id, current) index = json.loads( read_output_file(output, "annotation_studio/human_videos/index.json", current["token"]) ) except Exception as exc: raise HTTPException(400, "No human demo videos saved yet.") from exc media = f"/api/jobs/{job_id}/media?path=" return [ { key: entry.get(key) for key in ( "episode_index", "source_episode_index", "segment_index", "subtask", "human_task", "caption", "start_timestamp", "end_timestamp", "status", ) } | { kind: media + entry[f"{kind}_path"] for kind in ("robot_frame", "first_frame", "video") if HUMAN_VIDEO_FILE.fullmatch(entry.get(f"{kind}_path") or "") } for entry in index[:60] ] @app.get("/api/jobs/{job_id}/media") def human_video_media(job_id: str, path: str, request: Request): """Stream a generated video or frame from the user's (possibly private) output dataset.""" current = session(request, required=True) if not HUMAN_VIDEO_FILE.fullmatch(path): raise HTTPException(400, "Unsupported media path.") try: _, output = studio_output(job_id, current) content = read_output_file(output, path, current["token"]) except Exception as exc: raise HTTPException(404, "Media not found.") from exc return Response(content, media_type="video/mp4" if path.endswith(".mp4") else "image/png") @app.get("/health") def health(): return {"status": "ok", "service": "lerobot-annotation-studio"} @app.get("/api/models") def models(): return public_models() if (BASE / "frontend/dist/assets").exists(): app.mount("/assets", StaticFiles(directory=BASE / "frontend/dist/assets"), name="assets") @app.get("/{path:path}") def frontend(path: str): return FileResponse(BASE / "frontend/dist/index.html")