decision-studio / model_runtime_client.py
Xunzhuo's picture
Cursor
Serve HTTP pull workers from the built-in model runtime
c646299
Raw History Blame Contribute Delete
17.4 kB
"""Private client for the vLLM Semantic Router built-in model runtime.
`vllm-sr-runtime serve` (``vllm-sr serve <model>``) answers one state per
request on ``/v1/systemone`` and reports its artifact on ``/v1/models`` and in
each response's ``meta``, instead of the ``X-Decision-Artifact-*`` headers of
``vllm-sr decision serve``. This client keeps the same guarantees as
``DirectGateway``: every answer is bound to the pinned repository, revision,
manifest and model identity on the ``exact`` profile; the runtime's answers are
projected onto the strict SystemOne envelope without changing a value; and a
shared-question batch becomes one runtime request per state, reassembled in
request order.
"""
from __future__ import annotations
import asyncio
import json
import math
from collections.abc import Mapping
from types import MappingProxyType
import httpx
from direct_contract import (
CONFIDENCE_DEFINITIONS,
MARGIN_ORDINAL_V1,
validate_response_for_request,
)
from direct_gateway import (
MAX_BATCH_REQUEST_BYTES,
MAX_INFERENCE_RESPONSE_BYTES,
MAX_SINGLE_REQUEST_BYTES,
MAX_TIMEOUT_SECONDS,
MIN_TIMEOUT_SECONDS,
DirectGatewayError,
_hex_digest,
_loads_strict,
_private_origin,
_safe_retry_after,
)
SINGLE_PATH = "/v1/systemone"
HEALTH_PATH = "/health"
MODELS_PATH = "/v1/models"
MAX_CONTROL_RESPONSE_BYTES = 256 * 1024
MAX_CONTROL_TIMEOUT_SECONDS = 2.0
EXACT = "exact"
# Golden statuses of a runtime that verified its package and answers
# deterministically; a "failed" golden check never reaches readiness.
SERVING_GOLDEN = frozenset({"matched", "unverified"})
INPUT_ERRORS = frozenset({"invalid_question", "max_length_exceeded"})
RESPONSE_FIELDS = frozenset({"model", "answers", "usage", "meta"})
class ModelRuntimeGateway:
"""Attested client for one or more built-in runtime instances."""
def __init__(
self,
endpoints: Mapping[str, str],
*,
expected_artifacts: Mapping[str, Mapping[str, str]],
timeout_seconds: float = 45.0,
client: httpx.AsyncClient | None = None,
max_concurrency: int = 4,
):
expected = tuple(expected_artifacts)
if not expected or len(set(expected)) != len(expected):
raise ValueError("Configure a nonempty set of canonical Decision models")
if set(endpoints) != set(expected):
raise ValueError("runtime endpoints must cover exactly the configured models")
pinned = {}
confidence = {}
for model in expected:
item = expected_artifacts[model]
if not isinstance(item, Mapping):
raise ValueError("Pin each Decision model to a released artifact")
revision = item.get("revision")
manifest = item.get("manifest_sha256")
content = item.get("content_sha256")
if (
not _hex_digest(revision, 40)
or not _hex_digest(manifest, 64)
or (content is not None and not _hex_digest(content, 64))
):
raise ValueError("Pin each Decision model revision and manifest digest")
pinned[model] = MappingProxyType(
{
"revision": revision,
"manifest_sha256": manifest,
**({"content_sha256": content} if content is not None else {}),
}
)
confidence[model] = item.get("confidence_definition", MARGIN_ORDINAL_V1)
if confidence[model] not in CONFIDENCE_DEFINITIONS:
raise ValueError("Unknown Decision confidence definition")
if (
isinstance(timeout_seconds, bool)
or not isinstance(timeout_seconds, (int, float))
or not math.isfinite(timeout_seconds)
or not MIN_TIMEOUT_SECONDS <= timeout_seconds <= MAX_TIMEOUT_SECONDS
):
raise ValueError("runtime timeout must be between 0.1 and 300 seconds")
if type(max_concurrency) is not int or not 1 <= max_concurrency <= 64:
raise ValueError("Invalid runtime request concurrency")
normalized = {model: _private_origin(endpoints[model]) for model in expected}
if len(set(normalized.values())) != len(normalized):
raise ValueError("Each configured Decision model requires its own runtime origin")
self._endpoints = MappingProxyType(normalized)
self.expected_artifacts = MappingProxyType(pinned)
self.confidence_definitions = MappingProxyType(confidence)
self.models = frozenset(expected)
self.timeout_seconds = float(timeout_seconds)
self._max_concurrency = max_concurrency
self._owns_client = client is None
self._client = client or httpx.AsyncClient(
follow_redirects=False,
timeout=self.timeout_seconds,
trust_env=False,
)
async def aclose(self) -> None:
if self._owns_client:
await self._client.aclose()
# -- readiness -----------------------------------------------------------
async def ready(self, model: str) -> bool:
return (await self.probe(model))["loaded"]
async def probe(self, model: str) -> dict[str, object]:
"""Readiness plus the live artifact, with no private topology."""
if model not in self.models:
return {"loaded": False, "artifact": None}
timeout = min(self.timeout_seconds, MAX_CONTROL_TIMEOUT_SECONDS)
try:
health = await self._request_json("GET", model, HEALTH_PATH, timeout=timeout)
listing = await self._request_json("GET", model, MODELS_PATH, timeout=timeout)
except DirectGatewayError:
return {"loaded": False, "artifact": None}
artifact = self._card_artifact(model, listing)
loaded = (
isinstance(health, dict)
and health.get("status") == "ready"
and health.get("model") == model
and artifact is not None
and self._artifact_matches(model, artifact)
)
return {"loaded": bool(loaded), "artifact": artifact}
@staticmethod
def _card_artifact(model: str, listing: object) -> dict[str, str] | None:
if not isinstance(listing, dict) or not isinstance(listing.get("data"), list):
return None
if len(listing["data"]) != 1 or not isinstance(listing["data"][0], dict):
return None
card = listing["data"][0]
golden = card.get("golden")
if (
card.get("id") != model
or card.get("repo") != model
or card.get("ready") is not True
or card.get("profile") != EXACT
or not isinstance(golden, dict)
or golden.get("status") not in SERVING_GOLDEN
or not _hex_digest(card.get("revision"), 40)
or not _hex_digest(card.get("manifest_sha256"), 64)
or not _hex_digest(card.get("model_sha256"), 64)
):
return None
return {
"model": card["repo"],
"revision": card["revision"],
"manifest_sha256": card["manifest_sha256"],
"content_sha256": card["model_sha256"],
}
def _artifact_matches(self, model: str, artifact: Mapping[str, str]) -> bool:
pinned = self.expected_artifacts[model]
return (
artifact["model"] == model
and artifact["revision"] == pinned["revision"]
and artifact["manifest_sha256"] == pinned["manifest_sha256"]
and (
"content_sha256" not in pinned
or artifact["content_sha256"] == pinned["content_sha256"]
)
)
# -- inference -----------------------------------------------------------
async def evaluate(self, payload: Mapping[str, object], *, batch: bool) -> dict:
model = payload.get("model") if isinstance(payload, Mapping) else None
if not isinstance(model, str) or model not in self.models:
raise DirectGatewayError(422, "The selected canonical model is not configured")
if batch != ("states" in payload) or ("state" in payload and "states" in payload):
raise DirectGatewayError(422, "The Decision request does not match its route")
try:
encoded = json.dumps(
payload, ensure_ascii=False, separators=(",", ":"), allow_nan=False
).encode("utf-8")
except (TypeError, ValueError, UnicodeError, RecursionError) as exc:
raise DirectGatewayError(422, "The Decision request is not valid JSON") from exc
if len(encoded) > (MAX_BATCH_REQUEST_BYTES if batch else MAX_SINGLE_REQUEST_BYTES):
raise DirectGatewayError(413, "Request body exceeds the endpoint limit")
try:
async with asyncio.timeout(self.timeout_seconds):
if batch:
response = await self._batch(model, payload)
else:
response = await self._single(model, payload["state"], payload["questions"])
except TimeoutError as exc:
raise DirectGatewayError(
504, "The selected Decision runtime did not respond in time"
) from exc
try:
validate_response_for_request(
payload, response, batch=batch,
confidence=self.confidence_definitions[model],
)
except (TypeError, ValueError, KeyError, OverflowError, RecursionError) as exc:
raise DirectGatewayError(
502, "The selected Decision runtime returned an invalid response"
) from exc
return response
async def _batch(self, model: str, payload: Mapping[str, object]) -> dict:
states = payload["states"]
if not isinstance(states, list) or not states:
raise DirectGatewayError(422, "The Decision request does not match its route")
gate = asyncio.Semaphore(self._max_concurrency)
async def one(item):
if not isinstance(item, dict) or set(item) != {"id", "state"}:
raise DirectGatewayError(422, "Each context must contain exactly id and state")
async with gate:
return await self._single(model, item["state"], payload["questions"])
tasks = [asyncio.create_task(one(item)) for item in states]
try:
rows = await asyncio.gather(*tasks)
except BaseException:
for task in tasks:
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
raise
results = [
{"id": item["id"], "answers": row["answers"], "usage": row["usage"]}
for item, row in zip(states, rows, strict=True)
]
return {
"model": model,
"results": results,
"usage": {
"input_tokens": sum(row["usage"]["input_tokens"] for row in rows),
"output_tokens": sum(row["usage"]["output_tokens"] for row in rows),
},
}
async def _single(self, model: str, state: object, questions: object) -> dict:
body = {
"model": model,
"state": state,
"questions": questions,
"options": {"profile": EXACT, "return_meta": True},
}
response = await self._request_json(
"POST",
model,
SINGLE_PATH,
content=json.dumps(
body, ensure_ascii=False, separators=(",", ":"), allow_nan=False
).encode("utf-8"),
limit=MAX_INFERENCE_RESPONSE_BYTES,
)
if (
not isinstance(response, dict)
or set(response) != RESPONSE_FIELDS
or response["model"] != model
or not isinstance(response["answers"], dict)
or not isinstance(response["usage"], dict)
):
raise DirectGatewayError(
502, "The selected Decision runtime returned an invalid response"
)
if not self._meta_matches(model, response["meta"]):
raise DirectGatewayError(
503, "The selected Decision runtime artifact is unavailable"
)
errors = {
answer.get("error")
for answer in response["answers"].values()
if isinstance(answer, dict) and "error" in answer
}
if errors & INPUT_ERRORS:
raise DirectGatewayError(413, "Decision input exceeds the selected model limit")
if errors:
raise DirectGatewayError(503, "The selected Decision runtime is unavailable")
return {
"model": response["model"],
"answers": response["answers"],
"usage": response["usage"],
}
def _meta_matches(self, model: str, meta: object) -> bool:
pinned = self.expected_artifacts[model]
return (
isinstance(meta, dict)
and meta.get("revision") == pinned["revision"]
and meta.get("profile") == EXACT
and meta.get("numerics") == EXACT
and _hex_digest(meta.get("model_sha256"), 64)
and (
"content_sha256" not in pinned
or meta["model_sha256"] == pinned["content_sha256"]
)
)
async def _request_json(
self,
method: str,
model: str,
path: str,
*,
content: bytes | None = None,
limit: int = MAX_CONTROL_RESPONSE_BYTES,
timeout: float | None = None,
):
endpoint = self._endpoints[model]
request_timeout = timeout or self.timeout_seconds
headers = {"Accept": "application/json", "Accept-Encoding": "identity"}
if content is not None:
headers["Content-Type"] = "application/json"
try:
async with asyncio.timeout(request_timeout):
async with self._client.stream(
method,
endpoint + path,
content=content,
headers=headers,
follow_redirects=False,
timeout=request_timeout,
) as response:
if 300 <= response.status_code < 400:
raise DirectGatewayError(
502, "The selected Decision runtime returned an unsupported redirect"
)
if response.status_code >= 400 and not (
path == HEALTH_PATH and response.status_code == 503
):
self._raise_status(response)
return await self._read_json(response, limit)
except asyncio.CancelledError:
raise
except DirectGatewayError:
raise
except (TimeoutError, httpx.TimeoutException) as exc:
raise DirectGatewayError(
504, "The selected Decision runtime did not respond in time"
) from exc
except httpx.HTTPError as exc:
raise DirectGatewayError(503, "The selected Decision runtime is unavailable") from exc
@staticmethod
def _raise_status(response: httpx.Response) -> None:
status = response.status_code
if status in {400, 413}:
raise DirectGatewayError(413, "Decision input exceeds the selected model limit")
if status in {429, 529}:
raise DirectGatewayError(
status,
"The selected Decision runtime is temporarily overloaded",
retry_after=_safe_retry_after(response.headers.get("retry-after")),
)
if status == 503:
raise DirectGatewayError(
503,
"The selected Decision runtime is unavailable",
retry_after=_safe_retry_after(response.headers.get("retry-after")),
)
if status == 504:
raise DirectGatewayError(504, "The selected Decision runtime did not respond in time")
raise DirectGatewayError(502, "The selected Decision runtime rejected a validated request")
@staticmethod
async def _read_json(response: httpx.Response, limit: int):
content_type = response.headers.get("content-type", "").split(";", 1)[0].strip().lower()
if content_type != "application/json":
raise DirectGatewayError(
502, "The selected Decision runtime returned an invalid content type"
)
encoding = response.headers.get("content-encoding", "identity").strip().lower()
if encoding not in {"", "identity"}:
raise DirectGatewayError(
502, "The selected Decision runtime returned an unsupported encoding"
)
body = bytearray()
async for chunk in response.aiter_bytes(chunk_size=64 * 1024):
if len(chunk) > limit - len(body):
raise DirectGatewayError(
502, "The selected Decision runtime returned an oversized response"
)
body.extend(chunk)
try:
return _loads_strict(bytes(body))
except (TypeError, ValueError, UnicodeError, RecursionError) as exc:
raise DirectGatewayError(
502, "The selected Decision runtime returned invalid JSON"
) from exc