Spaces:
Running
Running
Download tests/test_generation_split.py from vllm-sr/decision-studio: direct link, hf CLI and curl.
- Browser
- Download file 15.3 kB
-
https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/tests/test_generation_split.py
- Command line
-
hf download hf://spaces/vllm-sr/decision-studio/tests/test_generation_split.py
-
curl -L -o test_generation_split.py https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/tests/test_generation_split.py
15.3 kB
| """A generation-only Studio cannot advertise or admit the other generation.""" | |
| import asyncio | |
| import os | |
| import unittest | |
| from contextlib import suppress | |
| from unittest.mock import patch | |
| import httpx | |
| from fastapi.testclient import TestClient | |
| from app import create_app | |
| from contract import MODEL | |
| from model_registry import PROFILES, model_registry | |
| from tetris_arena import HTTPDecisionAdapter, LOCAL_MODELS | |
| TOKEN = "t" * 32 | |
| QUESTIONS = {"accept": {"type": "noul", "instructions": "Is this accepted?"}} | |
| def generation_registry(version): | |
| return [ | |
| { | |
| "id": wire, | |
| "label": profile["repo_id"].rsplit("/", 1)[-1].split("-")[2], | |
| "version": version, | |
| "manifest_sha256": f"{index + 1:064x}", | |
| "revision": f"{index + 1:040x}", | |
| } | |
| for index, (wire, profile) in enumerate(PROFILES.items()) | |
| if f"/Decision-{version}-" in profile["repo_id"] | |
| ] | |
| class GenerationSplitTests(unittest.TestCase): | |
| def test_each_generation_is_an_independent_valid_registry(self): | |
| for version in ("1.0", "2.0"): | |
| with self.subTest(version=version): | |
| registry = model_registry(generation_registry(version)) | |
| self.assertEqual(len(registry), 6) | |
| self.assertTrue(all(f"/Decision-{version}-" in row["repo_id"] | |
| for row in registry.values())) | |
| def test_native_rejects_single_models_other_than_decision1_kai(self): | |
| for wire in ("decision2-kai", "decision-lux"): | |
| version = "2.0" if wire.startswith("decision2-") else "1.0" | |
| entry = next(row for row in generation_registry(version) if row["id"] == wire) | |
| with self.subTest(wire=wire), self.assertRaisesRegex(ValueError, "Native mode requires"): | |
| create_app(mode="native", registry=[entry]) | |
| def test_pull_queue_discovery_readiness_and_admission_follow_generation(self): | |
| for version, other in (("1.0", "2.0"), ("2.0", "1.0")): | |
| with self.subTest(version=version), patch.dict(os.environ, { | |
| "DECISION_WORKER_TOKEN": TOKEN, | |
| "TETRIS_LOCAL_API_URL": "http://127.0.0.1:7860/v1/systemone", | |
| "TETRIS_JEV_API_URL": "https://cloud.example.test/v1/systemone", | |
| "TETRIS_JEV_API_KEY": "test-cloud-key", | |
| }, clear=True): | |
| entries = generation_registry(version) | |
| api = create_app(mode="pull_queue", registry=entries) | |
| with TestClient(api) as client: | |
| expected = {PROFILES[row["id"]]["repo_id"] for row in entries} | |
| self.assertEqual(set(api.state.relays), {row["id"] for row in entries}) | |
| self.assertEqual({row["id"] for row in client.get("/v1/models").json()["models"]}, expected) | |
| self.assertEqual(client.get("/api/ready").status_code, 503) | |
| for index, row in enumerate(entries): | |
| heartbeat = client.post("/internal/worker/heartbeat", headers={ | |
| "Authorization": "Bearer " + TOKEN, | |
| }, json={ | |
| "model": row["id"], | |
| "manifest_sha256": row["manifest_sha256"], | |
| "worker_id": f"{index + 1:032x}", | |
| "phase": "ready", | |
| "capabilities": ["context_batch_v1"], | |
| }) | |
| self.assertEqual(heartbeat.status_code, 200, heartbeat.text) | |
| for data in ({"state": "A request"}, {"states": [{"id": "a", "state": "A request"}]}): | |
| job = client.post("/api/jobs", json={ | |
| "model": row["id"], "questions": QUESTIONS, **data, | |
| }) | |
| self.assertEqual(job.status_code, 202, job.text) | |
| self.assertEqual(job.json()["model"], row["id"]) | |
| client.delete(f"/api/jobs/{job.json()['id']}", params={"model": row["id"]}) | |
| ready = client.get("/api/ready") | |
| self.assertEqual(ready.status_code, 200) | |
| self.assertEqual(ready.json()["models"], dict.fromkeys(expected, True)) | |
| catalog = client.get("/api/tetris/config").json()["competitors"] | |
| allowed_competitors = {row.id for row in LOCAL_MODELS if row.request_model in expected} | |
| self.assertEqual({row["id"] for row in catalog if row["family"] == "decision"}, allowed_competitors) | |
| self.assertTrue(all(row["ready"] for row in catalog if row["family"] == "decision")) | |
| self.assertTrue(next(row for row in catalog if row["id"] == "jev-cloud")["ready"]) | |
| excluded = generation_registry(other)[0] | |
| canonical = PROFILES[excluded["id"]]["repo_id"] | |
| for route, data in (("/v1/systemone", {"state": "A request"}), | |
| ("/v1/systemone/batches", {"states": [{"id": "a", "state": "A request"}]})): | |
| self.assertEqual(client.post(route, json={ | |
| "model": canonical, "questions": QUESTIONS, **data, | |
| }).status_code, 422) | |
| self.assertEqual(client.post("/api/jobs", json={ | |
| "model": excluded["id"], "state": "A request", "questions": QUESTIONS, | |
| }).status_code, 422) | |
| excluded_competitor = next(row.id for row in LOCAL_MODELS if row.request_model == canonical) | |
| race = client.post("/api/tetris/races", json={ | |
| "left": excluded_competitor, | |
| "right": next(iter(allowed_competitors)), | |
| "mode": "steps", "max_steps": 1, | |
| }) | |
| self.assertEqual(race.status_code, 422, race.text) | |
| self.assertEqual(api.state.tetris.sessions, {}) | |
| def test_tetris_allowlist_rejects_unknown_model(self): | |
| with self.assertRaisesRegex(ValueError, "Unknown local Decision model allowlist"): | |
| HTTPDecisionAdapter.from_environment({}, allowed_local_models={"unknown/model"}) | |
| class StudioGenerationViewTests(unittest.TestCase): | |
| def test_generation_view_filters_discovery_and_arena_without_changing_queues(self): | |
| entries = generation_registry("2.0") + generation_registry("1.0") | |
| for version in ("", "2.0", "1.0"): | |
| with self.subTest(version=version), patch.dict(os.environ, { | |
| "DECISION_WORKER_TOKEN": TOKEN, | |
| "DECISION_STUDIO_GENERATION": version, | |
| "TETRIS_LOCAL_API_URL": "http://127.0.0.1:7860/v1/systemone", | |
| "TETRIS_JEV_API_URL": "https://jev.example.test/v1/systemone", | |
| "TETRIS_JEV_API_KEY": "test-upstream-secret", | |
| "TETRIS_SYSTEMONE_API_URL": "https://mirror.example.test/v1/systemone", | |
| "TETRIS_SYSTEMONE_API_KEY": "test-cloud-secret", | |
| }, clear=True): | |
| api = create_app(mode="pull_queue", registry=entries) | |
| expected = {row["id"] for row in entries | |
| if not version or row["version"] == version} | |
| with TestClient(api) as client: | |
| document = client.get("/v1/models").json() | |
| self.assertEqual({row["wire_id"] for row in document["models"]}, expected) | |
| self.assertEqual(document["default"], document["models"][0]["id"]) | |
| for row in document["models"]: | |
| configured = api.state.registry[row["wire_id"]] | |
| self.assertEqual(row["revision"], configured["revision"]) | |
| self.assertEqual(row["manifest_sha256"], configured["manifest_sha256"]) | |
| self.assertEqual(len(api.state.registry), 12) | |
| self.assertEqual(len(api.state.relays), 12) | |
| self.assertEqual(len(client.get("/api/ready").json()["models"]), 12) | |
| catalog = client.get("/api/tetris/config").json()["competitors"] | |
| expected_repos = {PROFILES[wire]["repo_id"] for wire in expected} | |
| self.assertEqual({row["id"] for row in catalog if row["family"] == "decision"}, | |
| {row.id for row in LOCAL_MODELS if row.request_model in expected_repos}) | |
| self.assertEqual({row["id"] for row in catalog if row["family"] == "cloud"}, | |
| set() if version == "2.0" else {"jev-cloud", "jev-cloud-mirror"}) | |
| self.assertNotIn("test-upstream-secret", repr(catalog)) | |
| self.assertNotIn("test-cloud-secret", repr(catalog)) | |
| if version == "2.0": | |
| self.assertTrue(all(row["family"] == "decision" for row in catalog)) | |
| self.assertNotIn("jev-cloud", api.state.tetris.adapter._endpoints) | |
| self.assertNotIn("jev-cloud-mirror", api.state.tetris.adapter._endpoints) | |
| for cloud in ("jev-cloud", "jev-cloud-mirror"): | |
| race = client.post("/api/tetris/races", json={ | |
| "left": cloud, "right": cloud, "mode": "steps", "max_steps": 1, | |
| }) | |
| self.assertEqual(race.status_code, 422, race.text) | |
| self.assertEqual(api.state.tetris.sessions, {}) | |
| else: | |
| self.assertTrue(all(row["ready"] for row in catalog if row["family"] == "cloud")) | |
| status = client.get("/api/status", params={"model": MODEL}) | |
| self.assertEqual(status.status_code, 200) | |
| self.assertEqual(status.json()["model"], MODEL) | |
| class SharedGatewayCompatibilityTests(unittest.IsolatedAsyncioTestCase): | |
| async def test_decision1_single_and_batch_still_complete_behind_decision2_view(self): | |
| entries = generation_registry("2.0") + generation_registry("1.0") | |
| entry = next(row for row in entries if row["id"] == MODEL) | |
| canonical = PROFILES[MODEL]["repo_id"] | |
| runtime = PROFILES[MODEL]["runtime_model"] | |
| worker_fields = { | |
| "model": MODEL, | |
| "manifest_sha256": entry["manifest_sha256"], | |
| "worker_id": "1" * 32, | |
| } | |
| headers = {"Authorization": "Bearer " + TOKEN} | |
| questions = { | |
| "route": {"type": "choice", "instructions": "Choose a route.", | |
| "criteria": {"allow": None, "deny": None}}, | |
| "accept": {"type": "noul", "instructions": "Is this accepted?"}, | |
| "urgency": {"type": "score", "instructions": "Rate urgency.", | |
| "criteria": ["Low", "Medium", "High"]}, | |
| } | |
| answers = { | |
| "route": {"type": "choice", "choice": "allow", "confidence": 0.5, | |
| "probabilities": {"allow": 0.75, "deny": 0.25}}, | |
| "accept": {"type": "noul", "noul": 0.8}, | |
| "urgency": {"type": "score", "score": 1.0, "confidence": 0.4, | |
| "probabilities": {"0": 0.2, "1": 0.6, "2": 0.2}, | |
| "legend": {"0": "Low", "1": "Medium", "2": "High"}}, | |
| } | |
| with patch.dict(os.environ, { | |
| "DECISION_WORKER_TOKEN": TOKEN, | |
| "DECISION_STUDIO_GENERATION": "2.0", | |
| }, clear=True): | |
| api = create_app(mode="pull_queue", registry=entries) | |
| try: | |
| async with httpx.AsyncClient(transport=httpx.ASGITransport(app=api), | |
| base_url="http://testserver") as client: | |
| heartbeat = await client.post("/internal/worker/heartbeat", headers=headers, | |
| json={**worker_fields, "phase": "ready", | |
| "capabilities": ["context_batch_v1"]}) | |
| self.assertEqual(heartbeat.status_code, 200) | |
| for batch in (False, True): | |
| with self.subTest(batch=batch): | |
| data = ({"states": [{"id": "a", "state": "One"}, | |
| {"id": "b", "state": "Two"}]} | |
| if batch else {"state": "One"}) | |
| route = "/v1/systemone/batches" if batch else "/v1/systemone" | |
| submitted = asyncio.create_task(client.post(route, json={ | |
| "model": canonical, "questions": questions, **data, | |
| })) | |
| try: | |
| claimed = await client.post("/internal/worker/claim", headers=headers, | |
| json={**worker_fields, "wait_seconds": 1}) | |
| self.assertEqual(claimed.status_code, 200) | |
| job = claimed.json()["job"] | |
| self.assertIsNotNone(job) | |
| self.assertEqual(job["body"]["model"], MODEL) | |
| response = { | |
| "model": runtime, | |
| "usage": {"input_tokens": 120 if batch else 60, "output_tokens": 0}, | |
| } | |
| if batch: | |
| response["results"] = [ | |
| {"id": row["id"], "answers": answers, | |
| "usage": {"input_tokens": 60, "output_tokens": 0}} | |
| for row in data["states"] | |
| ] | |
| else: | |
| response["answers"] = answers | |
| completed = await client.post("/internal/worker/result", headers=headers, | |
| json={**worker_fields, "id": job["id"], | |
| "lease_token": job["lease_token"], | |
| "result": {"kind": "http_runtime_v1", | |
| "response": response}}) | |
| self.assertEqual(completed.status_code, 200, completed.text) | |
| public = await asyncio.wait_for(submitted, timeout=3) | |
| self.assertEqual(public.status_code, 200, public.text) | |
| self.assertEqual(public.json(), dict(response, model=canonical)) | |
| finally: | |
| if not submitted.done(): | |
| submitted.cancel() | |
| with suppress(asyncio.CancelledError): | |
| await submitted | |
| finally: | |
| await api.state.tetris.aclose() | |
| if __name__ == "__main__": | |
| unittest.main() | |