decision-studio / http_pull_worker.py
Xunzhuo's picture
Cursor
Accept Decision 1.0 results under the runtime's attested Hub ID
59ba696
Raw History Blame Contribute Delete
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