Spaces:
Running
Running
Download server.py from lerobot/annotation-studio: direct link, hf CLI and curl.
- Browser
- Download file 18.8 kB
-
https://huggingface.co/spaces/lerobot/annotation-studio/resolve/main/server.py
- Command line
-
hf download hf://spaces/lerobot/annotation-studio/server.py
-
curl -L -o server.py https://huggingface.co/spaces/lerobot/annotation-studio/resolve/main/server.py
18.8 kB
| """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." | |
| 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 | |
| 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 | |
| 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) | |
| 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) | |
| 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 | |
| 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() | |
| 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://huggingface.co/proxy/lerobot-visualize-dataset.hf.space/{s['dataset']}/{s['episodes'][0]['episode_index']}", | |
| } | |
| def prompts(): | |
| return public_catalog() | |
| 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 | |
| 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}" | |
| 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 | |
| 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)" | |
| ) | |
| 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 | |
| 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] | |
| ] | |
| 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") | |
| def health(): | |
| return {"status": "ok", "service": "lerobot-annotation-studio"} | |
| def models(): | |
| return public_models() | |
| if (BASE / "frontend/dist/assets").exists(): | |
| app.mount("/assets", StaticFiles(directory=BASE / "frontend/dist/assets"), name="assets") | |
| def frontend(path: str): | |
| return FileResponse(BASE / "frontend/dist/index.html") | |