Spaces:
Sleeping
Sleeping
| """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 | |
| 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()) | |