Spaces:
Running
Running
File size: 7,157 Bytes
8827155 c646299 8827155 c646299 8827155 c646299 8827155 c646299 8827155 59ba696 8827155 c646299 8827155 c646299 8827155 c646299 8827155 82cafc8 8827155 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | """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
|