Spaces:
Running
Running
Download http_pull_worker.py from vllm-sr/decision-studio: direct link, hf CLI and curl.
- Browser
- Download file 7.16 kB
-
https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/http_pull_worker.py
- Command line
-
hf download hf://spaces/vllm-sr/decision-studio/http_pull_worker.py
-
curl -L -o http_pull_worker.py https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/http_pull_worker.py
7.16 kB
| """Outbound queue worker for a pinned private Decision HTTP runtime.""" | |
| import asyncio | |
| import concurrent.futures | |
| import json | |
| import logging | |
| import math | |
| import os | |
| import signal | |
| import threading | |
| from contract import to_records | |
| from direct_gateway import DirectGateway, DirectGatewayError | |
| from model_registry import PROFILES, model_registry | |
| from model_runtime_client import ModelRuntimeGateway | |
| from pull_worker import Gateway, Worker | |
| LOG = logging.getLogger("decision.http_worker") | |
| # `decision_serve_v1`: `vllm-sr decision serve` (artifact headers, native batches). | |
| # `vllm_sr_runtime_v1`: the built-in model runtime (`vllm-sr-runtime serve`). | |
| RUNTIME_APIS = { | |
| "decision_serve_v1": DirectGateway, | |
| "vllm_sr_runtime_v1": ModelRuntimeGateway, | |
| } | |
| class HTTPRuntime: | |
| """Run one attested private client on its own event loop for probe and inference.""" | |
| def __init__(self, model, origin, artifact, *, timeout_seconds=40, client=None, | |
| api="decision_serve_v1"): | |
| if model not in PROFILES or artifact.get("repo_id") != PROFILES[model]["repo_id"]: | |
| raise ValueError("Select one exact Decision model and canonical artifact") | |
| if api not in RUNTIME_APIS: | |
| raise ValueError("Select a supported Decision runtime API") | |
| if ( | |
| isinstance(timeout_seconds, bool) | |
| or not isinstance(timeout_seconds, (int, float)) | |
| or not math.isfinite(timeout_seconds) | |
| or not 0.1 <= timeout_seconds <= 90 | |
| ): | |
| raise ValueError("Runtime timeout must be between 0.1 and 90 seconds") | |
| self.model = model | |
| self.canonical_model = PROFILES[model].get("runtime_model", artifact["repo_id"]) | |
| self.manifest = artifact["manifest_sha256"] | |
| self.phase = "not_loaded" | |
| self._gateway = RUNTIME_APIS[api]( | |
| {self.canonical_model: origin}, | |
| expected_artifacts={self.canonical_model: { | |
| "revision": artifact["revision"], | |
| "manifest_sha256": self.manifest, | |
| **({"content_sha256": artifact["content_sha256"]} | |
| if "content_sha256" in artifact else {}), | |
| "confidence_definition": PROFILES[model]["confidence"], | |
| }}, | |
| timeout_seconds=timeout_seconds, | |
| client=client, | |
| ) | |
| self._loop = None | |
| self._thread = None | |
| def _submit(self, coroutine, timeout): | |
| if self._loop is None: | |
| coroutine.close() | |
| raise RuntimeError("Runtime client is not started") | |
| future = asyncio.run_coroutine_threadsafe(coroutine, self._loop) | |
| try: | |
| return future.result(timeout=timeout) | |
| except concurrent.futures.TimeoutError: | |
| future.cancel() | |
| raise TimeoutError("Runtime call exceeded its deadline") from None | |
| def _load(self): | |
| if self._loop is not None: | |
| return | |
| self._loop = asyncio.new_event_loop() | |
| def serve(): | |
| asyncio.set_event_loop(self._loop) | |
| self._loop.run_forever() | |
| self._thread = threading.Thread(target=serve, name="decision-runtime-http", daemon=True) | |
| self._thread.start() | |
| if not self.ready(): | |
| raise RuntimeError("Selected Decision runtime is not ready with its pinned artifact") | |
| self.phase = "ready" | |
| def ready(self): | |
| observed = self._submit(self._gateway.probe(self.canonical_model), 5) | |
| return observed["loaded"] is True | |
| def evaluate(self, body, records): | |
| if body.get("model") != self.model or records != to_records(body, model=self.model): | |
| raise ValueError("Request does not match this worker") | |
| request = dict(body, model=self.canonical_model) | |
| self.phase = "running" | |
| try: | |
| try: | |
| result = self._submit( | |
| self._gateway.evaluate(request, batch="states" in body), | |
| self._gateway.timeout_seconds + 5, | |
| ) | |
| except DirectGatewayError as exc: | |
| if exc.code == 413: | |
| raise ValueError("Selected model rejected the complete input") from None | |
| raise | |
| return {"kind": "http_runtime_v1", "response": result} | |
| finally: | |
| self.phase = "ready" | |
| def close(self): | |
| if self._loop is None: | |
| return | |
| try: | |
| self._submit(self._gateway.aclose(), 5) | |
| finally: | |
| self._loop.call_soon_threadsafe(self._loop.stop) | |
| self._thread.join(timeout=5) | |
| if self._thread.is_alive(): | |
| raise RuntimeError("Runtime client did not stop") | |
| self._loop.close() | |
| self._loop = None | |
| self._thread = None | |
| class HTTPWorker(Worker): | |
| heartbeat_join_timeout = 38 | |
| def heartbeat(self): | |
| if self.runtime.phase != "running" and not self.runtime.ready(): | |
| raise RuntimeError("Selected Decision runtime is unavailable") | |
| return super().heartbeat() | |
| def run(self): | |
| try: | |
| super().run() | |
| finally: | |
| self.runtime.close() | |
| def runtime_from_environment(): | |
| raw_registry = os.getenv("DECISION_MODEL_REGISTRY_V2", "") | |
| if not raw_registry: | |
| raise ValueError("DECISION_MODEL_REGISTRY_V2 is required") | |
| try: | |
| registry = model_registry(json.loads(raw_registry)) | |
| except (TypeError, ValueError, KeyError, RecursionError) as exc: | |
| raise ValueError("Invalid Decision model registry") from exc | |
| model = os.getenv("DECISION_WORKER_MODEL", "") | |
| if model not in registry: | |
| raise ValueError("DECISION_WORKER_MODEL must be an exact configured wire ID") | |
| item = registry[model] | |
| if "revision" not in item: | |
| raise ValueError("The selected model requires a pinned Hub revision") | |
| try: | |
| timeout = float(os.getenv("DECISION_RUNTIME_TIMEOUT_SECONDS", "40")) | |
| except ValueError as exc: | |
| raise ValueError("DECISION_RUNTIME_TIMEOUT_SECONDS must be numeric") from exc | |
| return HTTPRuntime(model, os.getenv("DECISION_RUNTIME_URL", ""), item, | |
| timeout_seconds=timeout, | |
| api=os.getenv("DECISION_RUNTIME_API", "decision_serve_v1")) | |
| def main(): | |
| logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") | |
| # Every heartbeat probes the runtime; keep per-request transport lines out of the log. | |
| logging.getLogger("httpx").setLevel(logging.WARNING) | |
| runtime = runtime_from_environment() | |
| gateway = Gateway(os.getenv("DECISION_STUDIO_URL", ""), | |
| os.getenv("DECISION_WORKER_TOKEN", ""), | |
| allow_loopback=os.getenv("DECISION_ALLOW_HTTP_LOOPBACK") == "1") | |
| worker = HTTPWorker(gateway, runtime) | |
| for sig in (signal.SIGINT, signal.SIGTERM): | |
| signal.signal(sig, lambda *_: worker.stop.set()) | |
| worker.run() | |
| if __name__ == "__main__": | |
| try: | |
| main() | |
| except Exception as exc: # noqa: BLE001 - never print secret-bearing exception text | |
| LOG.error("HTTP worker stopped: %s", type(exc).__name__) | |
| raise SystemExit(1) from None | |