zerogpu-transport-spike / validate_transport.py
giarom's picture
Update ZeroGPU transport spike
c43d35e verified
Raw
History Blame Contribute Delete
6.05 kB
"""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())