decision-studio / tests /test_model_runtime_client.py
Xunzhuo's picture
Cursor
Serve HTTP pull workers from the built-in model runtime
c646299
Raw History Blame Contribute Delete
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()