Spaces:
Running
Running
Download tests/test_model_runtime_client.py from vllm-sr/decision-studio: direct link, hf CLI and curl.
- Browser
- Download file 12.4 kB
-
https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/tests/test_model_runtime_client.py
- Command line
-
hf download hf://spaces/vllm-sr/decision-studio/tests/test_model_runtime_client.py
-
curl -L -o test_model_runtime_client.py https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/tests/test_model_runtime_client.py
12.4 kB
| """HTTP pull workers on the vLLM Semantic Router built-in model runtime.""" | |
| import asyncio | |
| import json | |
| import os | |
| import unittest | |
| from unittest.mock import patch | |
| import httpx | |
| from fastapi.testclient import TestClient | |
| from app import create_app | |
| from direct_contract import NORMALIZED_ENTROPY_V2 | |
| from direct_gateway import DirectGatewayError | |
| from http_pull_worker import HTTPRuntime, HTTPWorker, runtime_from_environment | |
| from model_registry import model_registry | |
| from model_runtime_client import ModelRuntimeGateway | |
| from pull_worker import GatewayError | |
| from relay import Relay | |
| from tests.test_decision2_registry import REGISTRY, questions, runtime_answers | |
| TOKEN = "t" * 32 | |
| WIRE = "decision2-lux" | |
| CANONICAL = "vllm-sr/Decision-2.0-Lux-9B" | |
| REVISION = "78bf3c03d9147aeb30b641edfe0e30ed04887ca5" | |
| MANIFEST = "4e0f4ca0cf3c833933f9bed3021012d061910016c9f4ca43b34f5d66e8d0a260" | |
| CONTENT = "0ece5faa210173f443f353474b419fd913a253b06645da3e8370c2db0e1c3339" | |
| ORIGIN = "http://172.31.231.15:8100" | |
| class FakeBuiltinRuntime: | |
| """The built-in runtime's HTTP contract: /health, /v1/models, /v1/systemone.""" | |
| def __init__(self): | |
| self.calls = [] | |
| self.status = "ready" | |
| self.card = { | |
| "id": CANONICAL, "object": "model", "repo": CANONICAL, "revision": REVISION, | |
| "model_sha256": CONTENT, "manifest_sha256": MANIFEST, "profile": "exact", | |
| "ready": True, "golden": {"status": "unverified", "checked": 0, "matched": 0}, | |
| } | |
| self.meta = {"revision": REVISION, "model_sha256": CONTENT, "profile": "exact", | |
| "numerics": "exact", "engine": "native", "accelerator": "rocm", | |
| "device": "rocm:0", "queue_ms": 0.1, "compute_ms": 9.0} | |
| self.question_error = None | |
| self.http_status = 200 | |
| def handler(self, request): | |
| if request.method == "GET" and request.url.path == "/health": | |
| return httpx.Response(200 if self.status == "ready" else 503, | |
| json={"status": self.status, "reason": None, "model": CANONICAL}) | |
| if request.method == "GET" and request.url.path == "/v1/models": | |
| return httpx.Response(200, json={"object": "list", "data": [dict(self.card)]}) | |
| body = json.loads(request.content) | |
| self.calls.append((request.url.path, body)) | |
| if self.http_status != 200: | |
| return httpx.Response(self.http_status, json={"error": {"code": "invalid_request", | |
| "message": "rejected"}}) | |
| answers = runtime_answers(body["state"]) | |
| if self.question_error: | |
| answers["domain"] = {"type": "choice", "error": self.question_error} | |
| return httpx.Response(200, json={ | |
| "model": CANONICAL, "answers": answers, | |
| "usage": {"input_tokens": 100 + len(body["state"]), "output_tokens": 0}, | |
| "meta": dict(self.meta), | |
| }) | |
| def gateway(runtime, **overrides): | |
| artifact = {"revision": REVISION, "manifest_sha256": MANIFEST, "content_sha256": CONTENT, | |
| "confidence_definition": NORMALIZED_ENTROPY_V2, **overrides} | |
| return ModelRuntimeGateway( | |
| {CANONICAL: ORIGIN}, expected_artifacts={CANONICAL: artifact}, | |
| client=httpx.AsyncClient(transport=httpx.MockTransport(runtime.handler)), | |
| ) | |
| class ModelRuntimeGatewayTests(unittest.TestCase): | |
| def setUp(self): | |
| self.runtime = FakeBuiltinRuntime() | |
| def run_async(self, coroutine): | |
| return asyncio.run(coroutine) | |
| def test_probe_requires_ready_exact_runtime_with_every_pin(self): | |
| self.assertTrue(self.run_async(gateway(self.runtime).ready(CANONICAL))) | |
| changes = ( | |
| ("revision", "f" * 40), ("manifest_sha256", "f" * 64), ("model_sha256", "f" * 64), | |
| ("repo", "vllm-sr/Decision-1.0-Lux-9B"), ("profile", "max_speed"), ("ready", False), | |
| ("golden", {"status": "failed"}), | |
| ) | |
| for field, value in changes: | |
| with self.subTest(field=field): | |
| runtime = FakeBuiltinRuntime() | |
| runtime.card[field] = value | |
| probe = self.run_async(gateway(runtime).probe(CANONICAL)) | |
| self.assertFalse(probe["loaded"]) | |
| loading = FakeBuiltinRuntime() | |
| loading.status = "loading" | |
| self.assertFalse(self.run_async(gateway(loading).ready(CANONICAL))) | |
| def test_single_answers_are_projected_unchanged(self): | |
| request = {"model": CANONICAL, "state": "Merge two sorted lists.", "questions": questions()} | |
| response = self.run_async(gateway(self.runtime).evaluate(request, batch=False)) | |
| self.assertEqual(response, { | |
| "model": CANONICAL, "answers": runtime_answers("Merge two sorted lists."), | |
| "usage": {"input_tokens": 123, "output_tokens": 0}, | |
| }) | |
| path, body = self.runtime.calls[-1] | |
| self.assertEqual(path, "/v1/systemone") | |
| self.assertEqual(body["options"], {"profile": "exact", "return_meta": True}) | |
| self.assertEqual(body["questions"], questions()) | |
| def test_batch_fans_out_per_state_in_request_order(self): | |
| states = [{"id": f"ctx-{i}", "state": "x" * (i + 1)} for i in range(9)] | |
| request = {"model": CANONICAL, "states": states, "questions": questions()} | |
| response = self.run_async(gateway(self.runtime).evaluate(request, batch=True)) | |
| self.assertEqual([row["id"] for row in response["results"]], [s["id"] for s in states]) | |
| for row, state in zip(response["results"], states): | |
| self.assertEqual(row["answers"], runtime_answers(state["state"])) | |
| self.assertEqual(response["usage"]["input_tokens"], | |
| sum(100 + len(s["state"]) for s in states)) | |
| self.assertEqual(sorted(body["state"] for _, body in self.runtime.calls), | |
| sorted(s["state"] for s in states)) | |
| def test_identity_errors_and_rejections_fail_closed(self): | |
| request = {"model": CANONICAL, "state": "x", "questions": questions()} | |
| cases = ( | |
| ("meta", {"revision": "f" * 40}, 503), | |
| ("meta", {"model_sha256": "f" * 64}, 503), | |
| ("meta", {"numerics": "approximate"}, 503), | |
| ("question_error", "max_length_exceeded", 413), | |
| ("question_error", "invalid_model_output", 503), | |
| ("http_status", 400, 413), | |
| ("http_status", 429, 429), | |
| ) | |
| for field, value, code in cases: | |
| with self.subTest(field=field, value=value): | |
| runtime = FakeBuiltinRuntime() | |
| if field == "meta": | |
| runtime.meta.update(value) | |
| else: | |
| setattr(runtime, field, value) | |
| with self.assertRaises(DirectGatewayError) as raised: | |
| self.run_async(gateway(runtime).evaluate(request, batch=False)) | |
| self.assertEqual(raised.exception.code, code) | |
| def test_origins_must_be_private_ip_literals(self): | |
| for origin in ("http://runtime-lux:8100", "http://8.8.8.8:8100", "http://172.31.231.15"): | |
| with self.subTest(origin=origin), self.assertRaises(ValueError): | |
| ModelRuntimeGateway({CANONICAL: origin}, expected_artifacts={ | |
| CANONICAL: {"revision": REVISION, "manifest_sha256": MANIFEST}}) | |
| class LocalGateway: | |
| def __init__(self, client): | |
| self.client = client | |
| def post(self, action, payload): | |
| response = self.client.post("/internal/worker/" + action, json=payload, | |
| headers={"Authorization": "Bearer " + TOKEN}) | |
| if response.status_code >= 400: | |
| raise GatewayError(response.status_code) | |
| return response.json() | |
| class BuiltinRuntimeWorkerTests(unittest.TestCase): | |
| """A Decision 2.0 worker on the built-in runtime, end to end through the queue.""" | |
| def setUp(self): | |
| self.runtime = FakeBuiltinRuntime() | |
| self.runtime_client = httpx.AsyncClient(transport=httpx.MockTransport(self.runtime.handler)) | |
| artifact = dict(next(item for item in REGISTRY if item["id"] == WIRE)) | |
| self.http_runtime = HTTPRuntime(WIRE, ORIGIN, artifact, client=self.runtime_client, | |
| api="vllm_sr_runtime_v1") | |
| self.app_context = TestClient(create_app( | |
| mode="pull_queue", | |
| relays={key: Relay(TOKEN, item["manifest_sha256"], model=key, | |
| complete_input_tokens=item["complete_input_tokens"]) | |
| for key, item in model_registry(REGISTRY).items()}, | |
| registry=REGISTRY, | |
| tetris_manager=type("NoTetris", (), {"public_config": lambda self: {}})(), | |
| )) | |
| self.client = self.app_context.__enter__() | |
| self.worker = HTTPWorker(LocalGateway(self.client), self.http_runtime) | |
| self.http_runtime._load() | |
| self.worker.heartbeat() | |
| def tearDown(self): | |
| self.http_runtime.close() | |
| asyncio.run(self.runtime_client.aclose()) | |
| self.app_context.__exit__(None, None, None) | |
| def test_studio_jobs_return_the_runtime_answers_unchanged(self): | |
| status = self.client.get("/api/status", params={"model": WIRE}).json() | |
| self.assertTrue(status["loaded"] and status["context_batch"]) | |
| self.assertEqual(status["complete_input_tokens"], 16384) | |
| for batch in (False, True): | |
| with self.subTest(batch=batch): | |
| request = {"model": WIRE, "questions": questions()} | |
| if batch: | |
| request["states"] = [{"id": "a", "state": "First"}, {"id": "b", "state": "Second!"}] | |
| else: | |
| request["state"] = "One request" | |
| submitted = self.client.post("/api/jobs", json=request) | |
| self.assertEqual(submitted.status_code, 202) | |
| claimed = self.worker.post("claim", {"worker_id": self.worker.worker_id, | |
| "wait_seconds": 0})["job"] | |
| self.worker.execute(claimed) | |
| result = self.client.get("/api/jobs/" + submitted.json()["id"], | |
| params={"model": WIRE}).json() | |
| self.assertEqual(result["status"], "succeeded") | |
| if batch: | |
| self.assertEqual([row["answers"] for row in result["result"]["results"]], | |
| [runtime_answers("First"), runtime_answers("Second!")]) | |
| else: | |
| self.assertEqual(result["result"]["answers"], runtime_answers("One request")) | |
| self.assertEqual(result["result"]["model"], CANONICAL) | |
| def test_space_rejects_a_decision1_confidence_on_a_decision2_queue(self): | |
| self.client.post("/api/jobs", json={"model": WIRE, "state": "x", "questions": questions()}) | |
| claimed = self.worker.post("claim", {"worker_id": self.worker.worker_id, | |
| "wait_seconds": 0})["job"] | |
| answers = runtime_answers("x") | |
| p = sorted(answers["domain"]["probabilities"].values()) | |
| answers["domain"]["confidence"] = p[-1] - p[-2] | |
| with self.assertRaises(GatewayError) as rejected: | |
| self.worker.post("result", { | |
| "worker_id": self.worker.worker_id, "id": claimed["id"], | |
| "lease_token": claimed["lease_token"], | |
| "result": {"kind": "http_runtime_v1", "response": { | |
| "model": CANONICAL, "answers": answers, | |
| "usage": {"input_tokens": 101, "output_tokens": 0}}}, | |
| }) | |
| self.assertEqual(rejected.exception.code, 422) | |
| def test_worker_environment_selects_the_runtime_api(self): | |
| environment = { | |
| "DECISION_MODEL_REGISTRY_V2": json.dumps(REGISTRY), | |
| "DECISION_WORKER_MODEL": WIRE, | |
| "DECISION_RUNTIME_URL": ORIGIN, | |
| } | |
| with patch.dict(os.environ, dict(environment, DECISION_RUNTIME_API="vllm_sr_runtime_v1")): | |
| self.assertIsInstance(runtime_from_environment()._gateway, ModelRuntimeGateway) | |
| with patch.dict(os.environ, dict(environment, DECISION_RUNTIME_API="vllm-sr-runtime")), \ | |
| self.assertRaisesRegex(ValueError, "runtime API"): | |
| runtime_from_environment() | |
| if __name__ == "__main__": | |
| unittest.main() | |