annotation-studio / server.py
pepijn223's picture
pepijn223 HF Staff
Add Human demo videos (robot → human, fal MiniMax-H3 480P) as a Video generation option (#1)
f19f993
Raw History Blame Contribute Delete
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."
@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://huggingface.co/proxy/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")