"""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