decision-studio / tests /test_http_pull_worker.py
Xunzhuo's picture
Separate Decision generations and introduce a cinematic Decision 2.0 studio
7222857 verified
Raw History Blame Contribute Delete
16.8 kB
"""The outbound HTTP worker carries attested, strict results through the queue."""
import asyncio
import io
import json
import os
import time
import unittest
from concurrent.futures import ThreadPoolExecutor
from unittest.mock import patch
import httpx
from fastapi.testclient import TestClient
from app import create_app
from contract import MODEL
from direct_gateway import ARTIFACT_RESPONSE_HEADERS
from http_pull_worker import HTTPRuntime, HTTPWorker, runtime_from_environment
from model_registry import MODEL_ORDER, PROFILES
from pull_worker import Gateway, GatewayError
from relay import Relay
TOKEN = "t" * 32
MANIFEST = "a" * 64
REVISION = "b" * 40
CONTENT = "c" * 64
CANONICAL = PROFILES[MODEL]["repo_id"]
# The deployed Decision 1.0 runtime attests its pre-move Hub ID.
ATTESTED = PROFILES[MODEL]["runtime_model"]
ORIGIN = "http://127.0.0.1:18401"
ARTIFACT = {
"repo_id": CANONICAL,
"manifest_sha256": MANIFEST,
"revision": REVISION,
"content_sha256": CONTENT,
}
class NoTetris:
def public_config(self):
return {"competitors": []}
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()
def payload(*, batch=False, model=MODEL):
request = {
"model": model,
"questions": {"route": {
"type": "choice",
"instructions": "Choose a route.",
"criteria": {"left": None, "right": None},
}},
}
if batch:
request["states"] = [
{"id": "one", "state": "First request"},
{"id": "two", "state": "Second request"},
]
else:
request["state"] = "One request"
return request
def answer():
return {
"type": "choice",
"choice": "left",
"confidence": 0.4,
"probabilities": {"left": 0.7, "right": 0.3},
}
def strict_response(request):
usage = {"input_tokens": 12, "output_tokens": 0}
if "states" in request:
return {
"model": request["model"],
"results": [
{"id": row["id"], "answers": {"route": answer()}, "usage": usage}
for row in request["states"]
],
"usage": {"input_tokens": 12 * len(request["states"]), "output_tokens": 0},
}
return {"model": request["model"], "answers": {"route": answer()}, "usage": usage}
class HTTPPullWorkerTests(unittest.TestCase):
def setUp(self):
self.calls = []
self.bad_header = False
self.bad_response = False
self.offline = False
def handler(request):
if request.method == "GET" and request.url.path == "/api/status":
return httpx.Response(200, json={
"status": "offline" if self.offline else "ready",
"artifact": {
"model": ATTESTED,
"revision": REVISION,
"manifest_sha256": MANIFEST,
"content_sha256": CONTENT,
},
})
body = json.loads(request.content)
self.calls.append((request.url.path, body))
response = strict_response(body)
if self.bad_response:
response["model"] = "another/model"
headers = {
header: {
"model": ATTESTED,
"revision": REVISION,
"manifest_sha256": "d" * 64 if self.bad_header else MANIFEST,
"content_sha256": CONTENT,
}[field]
for field, header in ARTIFACT_RESPONSE_HEADERS.items()
}
return httpx.Response(200, json=response, headers=headers)
self.runtime_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
self.runtime = HTTPRuntime(MODEL, ORIGIN, ARTIFACT, client=self.runtime_client)
self.app_context = TestClient(create_app(
mode="pull_queue",
relay=Relay(TOKEN, MANIFEST, model=MODEL),
registry=[{
"id": MODEL,
"label": "Kai",
"version": "1.0",
"manifest_sha256": MANIFEST,
}],
tetris_manager=NoTetris(),
))
self.client = self.app_context.__enter__()
self.worker = HTTPWorker(LocalGateway(self.client), self.runtime)
self.runtime._load()
self.worker.heartbeat()
def tearDown(self):
self.runtime.close()
asyncio.run(self.runtime_client.aclose())
self.app_context.__exit__(None, None, None)
def process_one(self):
claimed = self.worker.post("claim", {
"worker_id": self.worker.worker_id,
"wait_seconds": 0,
})["job"]
self.assertIsNotNone(claimed)
self.worker.execute(claimed)
return claimed
def test_single_and_batch_jobs_preserve_canonical_runtime_response(self):
for batch, expected_path in (
(False, "/v1/systemone"),
(True, "/v1/systemone/batches"),
):
with self.subTest(batch=batch):
submitted = self.client.post("/api/jobs", json=payload(batch=batch))
self.assertEqual(submitted.status_code, 202)
self.process_one()
completed = self.client.get("/api/jobs/" + submitted.json()["id"])
self.assertEqual(completed.status_code, 200)
self.assertEqual(completed.json()["status"], "succeeded")
self.assertEqual(completed.json()["result"],
strict_response(payload(batch=batch, model=CANONICAL)))
self.assertEqual(self.calls[-1],
(expected_path, payload(batch=batch, model=ATTESTED)))
def test_public_synchronous_route_returns_strict_result(self):
with ThreadPoolExecutor(max_workers=1) as executor:
pending = executor.submit(
self.client.post, "/v1/systemone", json=payload(model=CANONICAL)
)
for _ in range(200):
if self.client.get("/api/status").json()["queued"]:
break
time.sleep(0.01)
else:
self.fail("The public request was not queued")
self.process_one()
response = pending.result(timeout=5)
self.assertEqual(response.status_code, 200)
self.assertEqual(response.json(), strict_response(payload(model=CANONICAL)))
def test_studio_batch_routes_admit_over_256_kib_but_single_stays_bounded(self):
large_state = "x" * (300 * 1024)
batch_request = payload(batch=True)
batch_request["states"] = [{"id": "large", "state": large_state}]
single_request = payload()
single_request["state"] = large_state
for route in ("/api/jobs", "/api/evaluate"):
with self.subTest(route=route):
self.assertEqual(
self.client.post(route, json=single_request).status_code, 413
)
submitted = self.client.post("/api/jobs", json=batch_request)
self.assertEqual(submitted.status_code, 202)
self.process_one()
completed = self.client.get("/api/jobs/" + submitted.json()["id"]).json()
self.assertEqual(completed["status"], "succeeded")
with ThreadPoolExecutor(max_workers=1) as executor:
pending = executor.submit(
self.client.post, "/api/evaluate", json=batch_request
)
for _ in range(200):
if self.client.get("/api/status").json()["queued"]:
break
time.sleep(0.01)
else:
self.fail("The Studio batch was not queued")
self.process_one()
response = pending.result(timeout=5)
self.assertEqual(response.status_code, 200)
self.assertEqual(
response.json(), strict_response(dict(batch_request, model=CANONICAL))
)
def test_bad_artifact_or_response_fails_without_prediction(self):
for field in ("bad_header", "bad_response"):
with self.subTest(field=field):
setattr(self, field, True)
submitted = self.client.post("/api/jobs", json=payload())
self.assertEqual(submitted.status_code, 202)
self.process_one()
completed = self.client.get("/api/jobs/" + submitted.json()["id"]).json()
self.assertEqual(completed["status"], "failed")
self.assertNotIn("result", completed)
setattr(self, field, False)
def test_offline_runtime_does_not_send_a_fresh_heartbeat(self):
before = self.client.app.state.relay.worker["seen"]
self.offline = True
with self.assertRaises(RuntimeError):
self.worker.heartbeat()
self.assertEqual(self.client.app.state.relay.worker["seen"], before)
def test_space_rejects_mismatched_strict_result(self):
submitted = self.client.post("/api/jobs", json=payload()).json()
claimed = self.worker.post("claim", {
"worker_id": self.worker.worker_id,
"wait_seconds": 0,
})["job"]
wrong = strict_response(payload(model=ATTESTED))
wrong["model"] = "another/model"
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": wrong},
})
self.assertEqual(rejected.exception.code, 422)
self.worker.execute(claimed)
completed = self.client.get("/api/jobs/" + submitted["id"]).json()
self.assertEqual(completed["status"], "succeeded")
def test_strict_results_accept_reordered_question_maps(self):
questions = {
"beta": {"type": "noul", "instructions": "Is beta true?"},
"alpha": {"type": "noul", "instructions": "Is alpha true?"},
}
answers = {
"alpha": {"type": "noul", "noul": 0.2},
"beta": {"type": "noul", "noul": 0.8},
}
for batch in (False, True):
with self.subTest(batch=batch):
request = {"model": MODEL, "questions": questions}
if batch:
request["states"] = [
{"id": "first", "state": "First"},
{"id": "second", "state": "Second"},
]
strict = {
"model": ATTESTED,
"results": [
{"id": row["id"], "answers": answers,
"usage": {"input_tokens": 12, "output_tokens": 0}}
for row in request["states"]
],
"usage": {"input_tokens": 24, "output_tokens": 0},
}
else:
request["state"] = "One request"
strict = {
"model": ATTESTED,
"answers": answers,
"usage": {"input_tokens": 12, "output_tokens": 0},
}
submitted = self.client.post("/api/jobs", json=request).json()
claimed = self.worker.post("claim", {
"worker_id": self.worker.worker_id,
"wait_seconds": 0,
})["job"]
self.worker.post("result", {
"worker_id": self.worker.worker_id,
"id": claimed["id"],
"lease_token": claimed["lease_token"],
"result": {"kind": "http_runtime_v1", "response": strict},
})
completed = self.client.get("/api/jobs/" + submitted["id"]).json()
self.assertEqual(completed["status"], "succeeded")
self.assertEqual(completed["result"], dict(strict, model=CANONICAL))
class WorkerConfigurationTests(unittest.TestCase):
def test_claim_transport_accepts_a_large_valid_batch_body(self):
gateway = Gateway("https://example.test", TOKEN)
reply = json.dumps({"job": {"body": {"state": "x" * (1200 * 1024)}}}).encode()
class Opener:
def open(self, request, timeout):
self.request = request
return io.BytesIO(reply)
gateway.opener = Opener()
response = gateway.post("claim", {"worker_id": "1" * 32})
self.assertEqual(len(response["job"]["body"]["state"]), 1200 * 1024)
self.assertTrue(gateway.opener.request.full_url.endswith("/internal/worker/claim"))
def test_every_wire_id_binds_to_its_exact_canonical_model(self):
registry = [{
"id": model,
"label": model,
"version": "1.0",
"manifest_sha256": format(index + 1, "x") * 64,
"revision": format(index + 1, "x") * 40,
} for index, model in enumerate(MODEL_ORDER)]
for model in MODEL_ORDER:
with self.subTest(model=model), patch.dict(os.environ, {
"DECISION_MODEL_REGISTRY_V2": json.dumps(registry),
"DECISION_WORKER_MODEL": model,
"DECISION_RUNTIME_URL": ORIGIN,
}):
runtime = runtime_from_environment()
self.assertEqual(runtime.model, model)
self.assertEqual(runtime.canonical_model, PROFILES[model].get(
"runtime_model", PROFILES[model]["repo_id"]))
self.assertEqual(runtime.manifest, registry[MODEL_ORDER.index(model)]["manifest_sha256"])
with patch.dict(os.environ, {
"DECISION_MODEL_REGISTRY_V2": json.dumps(registry),
"DECISION_WORKER_MODEL": "decision-nano-preview",
"DECISION_RUNTIME_URL": ORIGIN,
}), self.assertRaises(ValueError):
runtime_from_environment()
class PullQueueTetrisReadinessTests(unittest.TestCase):
def test_offline_queue_is_unavailable_in_config_and_race_admission(self):
lux_manifest = "d" * 64
relays = {
MODEL: Relay(TOKEN, MANIFEST, model=MODEL),
"decision-lux": Relay(
TOKEN, lux_manifest, model="decision-lux",
complete_input_tokens=16384,
),
}
registry = [
{"id": MODEL, "label": "Kai", "version": "1.0",
"manifest_sha256": MANIFEST},
{"id": "decision-lux", "label": "Lux", "version": "1.0",
"manifest_sha256": lux_manifest},
]
with patch.dict(os.environ, {
"TETRIS_LOCAL_API_URL": "http://127.0.0.1:7860/v1/systemone",
}), TestClient(create_app(
mode="pull_queue", relays=relays, registry=registry,
)) as client:
headers = {"Authorization": "Bearer " + TOKEN}
ready = client.post("/internal/worker/heartbeat", json={
"model": MODEL,
"manifest_sha256": MANIFEST,
"worker_id": "1" * 32,
"phase": "ready",
"capabilities": ["context_batch_v1"],
}, headers=headers)
self.assertEqual(ready.status_code, 200)
config = client.get("/api/tetris/config").json()
by_id = {item["id"]: item["ready"] for item in config["competitors"]}
self.assertTrue(by_id["kai"])
self.assertFalse(by_id["lux"])
self.assertNotIn("nox", by_id)
race = {"left": "kai", "right": "lux", "mode": "steps", "max_steps": 1}
self.assertEqual(client.post("/api/tetris/races", json=race).status_code, 503)
ready = client.post("/internal/worker/heartbeat", json={
"model": "decision-lux",
"manifest_sha256": lux_manifest,
"worker_id": "2" * 32,
"phase": "ready",
"capabilities": ["context_batch_v1"],
}, headers=headers)
self.assertEqual(ready.status_code, 200)
config = client.get("/api/tetris/config").json()
by_id = {item["id"]: item["ready"] for item in config["competitors"]}
self.assertTrue(by_id["lux"])
if __name__ == "__main__":
unittest.main()