Spaces:
Running
Running
Download model_runtime_client.py from vllm-sr/decision-studio: direct link, hf CLI and curl.
- Browser
- Download file 17.4 kB
-
https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/model_runtime_client.py
- Command line
-
hf download hf://spaces/vllm-sr/decision-studio/model_runtime_client.py
-
curl -L -o model_runtime_client.py https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/model_runtime_client.py
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} | |
| 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 | |
| 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") | |
| 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 | |