"""Validate the deployed ZeroGPU transport spike. Examples: python spikes/zerogpu_transport/validate_transport.py https://giarom-zerogpu-transport-spike.hf.space --mode rest python spikes/zerogpu_transport/validate_transport.py https://giarom-zerogpu-transport-spike.hf.space --mode gradio """ from __future__ import annotations import argparse import json import queue import threading import time from dataclasses import dataclass from typing import Any from urllib.error import HTTPError from urllib.request import Request, urlopen @dataclass class CheckResult: name: str ok: bool detail: str = "" def _url(base_url: str, path: str) -> str: return f"{base_url.rstrip('/')}{path}" def _request_json(base_url: str, method: str, path: str, body: dict[str, Any] | None = None) -> dict[str, Any]: data = None if body is None else json.dumps(body).encode("utf-8") req = Request( _url(base_url, path), data=data, method=method, headers={"content-type": "application/json"}, ) try: with urlopen(req, timeout=30) as res: raw = res.read().decode("utf-8") return json.loads(raw) if raw else {} except HTTPError as exc: raw = exc.read().decode("utf-8") try: payload = json.loads(raw) except json.JSONDecodeError: payload = {"error": {"code": "HTTP_ERROR", "message": raw}} payload["_status"] = exc.code return payload def _read_sse(base_url: str, job_id: str, out: "queue.Queue[dict[str, Any]]", stop: threading.Event) -> None: req = Request(_url(base_url, f"/v1/jobs/{job_id}/events"), method="GET") with urlopen(req, timeout=60) as res: event_lines: list[str] = [] while not stop.is_set(): line = res.readline() if not line: break text = line.decode("utf-8").rstrip("\n") if text == "": for item in event_lines: if item.startswith("data: "): out.put(json.loads(item[6:])) event_lines = [] continue event_lines.append(text) def _run_page_rest(base_url: str, job_id: str, page_number: int) -> dict[str, Any]: return _request_json(base_url, "POST", f"/v1/jobs/{job_id}/pages/{page_number}/run", {}) def _run_page_gradio(base_url: str, job_id: str, page_number: int) -> dict[str, Any]: try: from gradio_client import Client except ImportError as exc: raise RuntimeError("gradio_client is required for --mode gradio") from exc client = Client(base_url.rstrip("/")) return client.predict(job_id=job_id, page_number=page_number, api_name="/run_page") def _check(condition: bool, name: str, detail: str = "") -> CheckResult: return CheckResult(name=name, ok=condition, detail=detail) def run_validation(base_url: str, mode: str, pages: int) -> list[CheckResult]: results: list[CheckResult] = [] live = _request_json(base_url, "GET", "/v1/health/live") results.append(_check(live.get("ok") is True, "health live", str(live))) ready = _request_json(base_url, "GET", "/v1/health/ready") results.append(_check(ready.get("ok") is True, "health ready", str(ready))) created = _request_json(base_url, "POST", "/v1/jobs", {"pages": pages}) job_id = created.get("job_id") results.append(_check(isinstance(job_id, str), "create job", str(created))) if not isinstance(job_id, str): return results events: "queue.Queue[dict[str, Any]]" = queue.Queue() stop = threading.Event() reader = threading.Thread(target=_read_sse, args=(base_url, job_id, events, stop), daemon=True) reader.start() time.sleep(0.5) runner = _run_page_rest if mode == "rest" else _run_page_gradio for page_number in range(1, pages + 1): page_result = runner(base_url, job_id, page_number) results.append(_check(page_result.get("ok") is True, f"{mode} page {page_number}", str(page_result))) deadline = time.time() + 20 seen: list[dict[str, Any]] = [] while time.time() < deadline: try: seen.append(events.get(timeout=0.5)) except queue.Empty: pass if any(event.get("event") == "job.completed" for event in seen): break stop.set() names = [event.get("event") for event in seen] results.append(_check("page.started" in names, "SSE page.started", str(names))) results.append(_check("page.completed" in names, "SSE page.completed", str(names))) results.append(_check("job.completed" in names, "SSE job.completed", str(names))) final_job = _request_json(base_url, "GET", f"/v1/jobs/{job_id}") results.append(_check(final_job.get("status") == "completed", "final job status", str(final_job))) cancelled_job = _request_json(base_url, "POST", "/v1/jobs", {"pages": 1}) cancel_id = cancelled_job.get("job_id") if isinstance(cancel_id, str): deleted = _request_json(base_url, "DELETE", f"/v1/jobs/{cancel_id}") results.append(_check(deleted.get("cancelled") is True, "DELETE cancellation", str(deleted))) cancelled_page = runner(base_url, cancel_id, 1) error_code = (cancelled_page.get("error") or {}).get("code") results.append(_check(error_code == "JOB_CANCELLED", "cancelled page rejected", str(cancelled_page))) return results def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("base_url") parser.add_argument("--mode", choices=["rest", "gradio"], default="rest") parser.add_argument("--pages", type=int, default=3) args = parser.parse_args() results = run_validation(args.base_url, args.mode, args.pages) failed = [result for result in results if not result.ok] for result in results: status = "PASS" if result.ok else "FAIL" print(f"{status} {result.name}: {result.detail}") return 1 if failed else 0 if __name__ == "__main__": raise SystemExit(main())