Contract, service and tests (fake engine, no GPU). Scores through the checkpoint's own inference.py (DecisionEngine.predict, sha256-pinned); maps semif decisions onto Jev choice questions, packs /decide/shared into calls of at most 16, runs orderings in waves, and keeps semif's error mapping, admission, body limit and hard VRAM cap. Deltas from semif-serve are listed in the contract.
435 lines
20 KiB
Python
435 lines
20 KiB
Python
"""intern-decision-serve HTTP behaviour against a fake engine (no torch, no model).
|
|
Contract: services/intern-decision-serve/intern-decision-serve.contract.md"""
|
|
import json
|
|
|
|
from fastapi.testclient import TestClient
|
|
|
|
from fake_engine import CALIBRATION, FakeEngine, request_sha
|
|
from intern_decision_serve.app import create_app
|
|
from intern_decision_serve.config import Settings
|
|
|
|
TOKEN = "t" * 40
|
|
AUTH = {"Authorization": f"Bearer {TOKEN}"}
|
|
OPTIONS = [{"id": "yes", "description": "It passed."}, {"id": "no", "description": "It did not pass at all."}]
|
|
ROW = {"id": "r1", "state": "The deploy passed.", "question": "Did it pass?", "options": OPTIONS}
|
|
|
|
|
|
def make_client(engine=None, **overrides):
|
|
return TestClient(create_app(Settings(api_token=TOKEN, **overrides), engine or FakeEngine()))
|
|
|
|
|
|
def test_decide_asks_one_question_named_q_and_maps_the_answer_back_in_option_order():
|
|
engine = FakeEngine()
|
|
response = make_client(engine).post("/decide", json=ROW, headers=AUTH)
|
|
assert response.status_code == 200
|
|
assert engine.calls == [{"state": "The deploy passed.", "questions": {"q": {
|
|
"type": "choice", "instructions": "Did it pass?",
|
|
"criteria": {"yes": "It passed.", "no": "It did not pass at all."}}}}]
|
|
out = response.json()
|
|
native, _ = FakeEngine().predict(engine.calls[0])
|
|
answer = native["answers"]["q"]
|
|
assert out["id"] == "r1"
|
|
assert out["option_ids"] == ["yes", "no"]
|
|
assert out["probabilities"] == [answer["probabilities"]["yes"], answer["probabilities"]["no"]]
|
|
assert out["top"] == "no" == answer["choice"]
|
|
assert out["confidence"] == answer["confidence"]
|
|
assert out["calibration"] == CALIBRATION
|
|
assert out["native"] == answer
|
|
assert out["input_tokens"] == 100
|
|
assert out["prompt_sha256"] == request_sha(engine.calls[0])
|
|
assert out["call"] == {"index": 0, "field": "q", "questions": 1}
|
|
assert out["forward_seconds"] == 0.0125
|
|
assert out["total_seconds"] >= 0
|
|
assert out["model"] == FakeEngine.metadata
|
|
for key in ("prompt_version", "readout", "probability_status"):
|
|
assert out[key]
|
|
|
|
|
|
import pytest # noqa: E402
|
|
|
|
|
|
@pytest.mark.parametrize("headers", [{}, {"Authorization": "Bearer wrong"}, {"Authorization": TOKEN}])
|
|
def test_posts_without_the_right_bearer_are_401_and_never_reach_the_engine(headers):
|
|
engine = FakeEngine()
|
|
client = make_client(engine)
|
|
for path, body in (("/decide", ROW), ("/decide/shared", {"state": "s", "decisions": [ROW]})):
|
|
response = client.post(path, json=body, headers=headers)
|
|
assert response.status_code == 401
|
|
assert response.json()["error"]["code"] == "unauthorized"
|
|
assert engine.calls == []
|
|
|
|
|
|
def test_health_needs_no_auth_and_reports_the_limits_and_the_chunking_rule():
|
|
response = make_client(vram_cap_gib=11.0).get("/health")
|
|
assert response.status_code == 200
|
|
out = response.json()
|
|
assert out["status"] == "ok"
|
|
assert out["model"] == FakeEngine().health()
|
|
assert (out["vram_cap_gib"], out["max_tokens"], out["max_decisions"]) == (11.0, 8192, 64)
|
|
assert out["max_questions_per_call"] == 16
|
|
assert "16" in out["chunking"]
|
|
assert out["workloads"] == []
|
|
|
|
|
|
def decisions(n, options=OPTIONS):
|
|
return [{"id": f"decision-{i}", "question": f"Question {i}?", "options": options} for i in range(n)]
|
|
|
|
|
|
def test_shared_asks_q1_to_qn_over_the_shared_state_in_one_call():
|
|
engine = FakeEngine()
|
|
body = {"state": {"deploy": "passed"}, "decisions": decisions(3)}
|
|
response = make_client(engine).post("/decide/shared", json=body, headers=AUTH)
|
|
assert response.status_code == 200
|
|
assert engine.calls == [{"state": {"deploy": "passed"}, "questions": {
|
|
f"q{i + 1}": {"type": "choice", "instructions": f"Question {i}?",
|
|
"criteria": {"yes": "It passed.", "no": "It did not pass at all."}} for i in range(3)}}]
|
|
out = response.json()
|
|
assert [r["id"] for r in out["results"]] == ["decision-0", "decision-1", "decision-2"]
|
|
assert [r["call"] for r in out["results"]] == [{"index": 0, "field": f"q{i}", "questions": 3} for i in (1, 2, 3)]
|
|
assert out["timing"]["calls"] == 1
|
|
assert out["timing"]["batch_size"] == 3
|
|
assert out["timing"]["questions_per_call"] == [3]
|
|
assert out["timing"]["input_tokens"] == [300]
|
|
assert out["timing"]["inference_seconds"] == 0.0125
|
|
assert out["timing"]["total_seconds"] >= 0
|
|
|
|
|
|
def test_one_shared_decision_is_the_same_call_as_decide():
|
|
engine = FakeEngine()
|
|
client = make_client(engine)
|
|
client.post("/decide", json=ROW, headers=AUTH)
|
|
client.post("/decide/shared", json={"state": ROW["state"], "decisions": [ROW]}, headers=AUTH)
|
|
assert engine.calls[0] == engine.calls[1]
|
|
|
|
|
|
@pytest.mark.parametrize("n, sizes", [(16, [16]), (17, [16, 1]), (40, [16, 16, 8])])
|
|
def test_more_than_16_questions_are_split_greedily_in_request_order(n, sizes):
|
|
engine = FakeEngine()
|
|
out = make_client(engine).post("/decide/shared", json={"state": "s", "decisions": decisions(n)},
|
|
headers=AUTH).json()
|
|
assert [len(c["questions"]) for c in engine.calls] == sizes
|
|
asked = [q["instructions"] for c in engine.calls for q in c["questions"].values()]
|
|
assert asked == [f"Question {i}?" for i in range(n)]
|
|
assert [r["id"] for r in out["results"]] == [f"decision-{i}" for i in range(n)]
|
|
for i, result in enumerate(out["results"]):
|
|
call = i // 16
|
|
size = sizes[call]
|
|
field = "q" if size == 1 else f"q{i % 16 + 1}"
|
|
assert result["call"] == {"index": call, "field": field, "questions": size}
|
|
assert result["prompt_sha256"] == request_sha(engine.calls[call])
|
|
assert out["timing"]["questions_per_call"] == sizes
|
|
assert out["timing"]["calls"] == len(sizes)
|
|
|
|
|
|
def test_decision_ids_never_reach_the_model():
|
|
engine = FakeEngine()
|
|
make_client(engine).post("/decide/shared", json={"state": "s", "decisions": decisions(20)}, headers=AUTH)
|
|
make_client(engine).post("/decide", json=ROW, headers=AUTH)
|
|
assert "decision-" not in repr(engine.calls) and "r1" not in repr(engine.calls)
|
|
|
|
|
|
from intern_decision_serve.errors import OutOfMemory, ScoringFailed # noqa: E402
|
|
|
|
|
|
class RaisingEngine(FakeEngine):
|
|
def __init__(self, exc, after=0):
|
|
super().__init__()
|
|
self.exc, self.after = exc, after
|
|
|
|
def predict(self, request):
|
|
if len(self.calls) >= self.after:
|
|
self.calls.append(request)
|
|
raise self.exc
|
|
return super().predict(request)
|
|
|
|
|
|
@pytest.mark.parametrize("exc, status, code", [
|
|
(ValueError("Example has 9000 tokens, above 8192; truncation is forbidden"), 422, "invalid_request"),
|
|
(OutOfMemory("CUDA out of memory. Tried to allocate 2.00 GiB."), 503, "out_of_memory"),
|
|
(ScoringFailed("RuntimeError: boom"), 500, "scoring_failed"),
|
|
(ArithmeticError("Floating-point temperature scaling changed argmax"), 500, "scoring_failed"),
|
|
])
|
|
def test_engine_failures_map_to_the_contract_codes_on_both_endpoints(exc, status, code):
|
|
for path, body in (("/decide", ROW), ("/decide/shared", {"state": "s", "decisions": decisions(2)})):
|
|
response = make_client(RaisingEngine(exc)).post(path, json=body, headers=AUTH)
|
|
assert response.status_code == status
|
|
assert response.json()["error"]["code"] == code
|
|
assert str(exc) in response.json()["error"]["message"]
|
|
|
|
|
|
def test_an_oom_in_a_later_chunk_fails_the_whole_request_with_503():
|
|
engine = RaisingEngine(OutOfMemory("CUDA out of memory."), after=1)
|
|
response = make_client(engine).post("/decide/shared", json={"state": "s", "decisions": decisions(20)},
|
|
headers=AUTH)
|
|
assert response.status_code == 503
|
|
assert len(engine.calls) == 2
|
|
|
|
|
|
class DroppingEngine(FakeEngine):
|
|
"""Returns an answer that lacks one of the asked option ids."""
|
|
def predict(self, request):
|
|
response, sha = super().predict(request)
|
|
for answer in response["answers"].values():
|
|
answer["probabilities"].pop("no")
|
|
return response, sha
|
|
|
|
|
|
def test_an_answer_missing_an_option_id_is_a_500_not_a_guess():
|
|
response = make_client(DroppingEngine()).post("/decide", json=ROW, headers=AUTH)
|
|
assert response.status_code == 500
|
|
assert response.json()["error"]["code"] == "scoring_failed"
|
|
assert "no" in response.json()["error"]["message"]
|
|
|
|
|
|
def opts(n):
|
|
return [{"id": f"o{i}", "description": f"Option {i}"} for i in range(n)]
|
|
|
|
|
|
BAD_DECIDE = [
|
|
{**ROW, "options": opts(1)},
|
|
{**ROW, "options": opts(17)},
|
|
{**ROW, "options": [{"id": "a", "description": "x"}, {"id": "a", "description": "y"}]},
|
|
{**ROW, "options": [{"id": "a"}, {"id": "b", "description": "y"}]},
|
|
{**ROW, "options": [{"id": 1, "description": "x"}, {"id": "b", "description": "y"}]},
|
|
{**ROW, "id": ""},
|
|
{**ROW, "question": ""},
|
|
{**ROW, "state": ""},
|
|
{**ROW, "state": {}},
|
|
{**ROW, "state": []},
|
|
{**ROW, "state": 7},
|
|
{**ROW, "workload": "triage"},
|
|
{k: v for k, v in ROW.items() if k != "question"},
|
|
]
|
|
|
|
|
|
@pytest.mark.parametrize("body", BAD_DECIDE)
|
|
def test_semif_rule_violations_are_422_before_the_engine_runs(body):
|
|
engine = FakeEngine()
|
|
client = make_client(engine)
|
|
response = client.post("/decide", json=body, headers=AUTH)
|
|
assert response.status_code == 422
|
|
assert response.json()["error"]["code"] == "invalid_request"
|
|
decision = {k: v for k, v in body.items() if k not in ("state", "workload")}
|
|
shared = {"state": body.get("state", "s"), "decisions": [decision],
|
|
**({"workload": body["workload"]} if "workload" in body else {})}
|
|
assert client.post("/decide/shared", json=shared, headers=AUTH).status_code == 422
|
|
assert engine.calls == []
|
|
|
|
|
|
def test_a_non_finite_state_is_422():
|
|
engine = FakeEngine()
|
|
raw = '{"id": "r", "state": {"x": NaN}, "question": "q?", "options": [{"id": "a", "description": "A"}, {"id": "b", "description": "B"}]}'
|
|
response = make_client(engine).post("/decide", content=raw, headers={**AUTH, "Content-Type": "application/json"})
|
|
assert response.status_code == 422
|
|
assert engine.calls == []
|
|
|
|
|
|
def test_duplicate_decision_ids_are_422():
|
|
body = {"state": "s", "decisions": [{**decisions(1)[0]}, {**decisions(1)[0]}]}
|
|
assert make_client().post("/decide/shared", json=body, headers=AUTH).status_code == 422
|
|
|
|
|
|
@pytest.mark.parametrize("raw", ["{not json", "[]", '{"state": "s"}', '{"state": "s", "decisions": "x"}'])
|
|
def test_malformed_bodies_are_422(raw):
|
|
client = make_client()
|
|
for path in ("/decide", "/decide/shared"):
|
|
response = client.post(path, content=raw, headers={**AUTH, "Content-Type": "application/json"})
|
|
assert response.status_code == 422
|
|
assert response.json()["error"]["code"] == "invalid_request"
|
|
|
|
|
|
@pytest.mark.parametrize("n", [0, 5])
|
|
def test_a_row_count_outside_1_to_max_decisions_is_422(n):
|
|
engine = FakeEngine()
|
|
response = make_client(engine, max_decisions=4).post(
|
|
"/decide/shared", json={"state": "s", "decisions": decisions(n)}, headers=AUTH)
|
|
assert response.status_code == 422
|
|
assert engine.calls == []
|
|
|
|
|
|
import math # noqa: E402
|
|
|
|
|
|
def position_bias(field, question):
|
|
"""Pure position bias: the first-listed option gets +2, whatever it says."""
|
|
return [2.0 if i == 0 else 0.0 for i in range(len(question["criteria"]))]
|
|
|
|
|
|
def test_rotations_ask_each_ordering_in_its_own_wave_and_cancel_a_position_bias_exactly():
|
|
engine = FakeEngine(scorer=position_bias)
|
|
body = {**ROW, "options": opts(3), "orderings": "rotations"}
|
|
out = make_client(engine).post("/decide", json=body, headers=AUTH).json()
|
|
assert [list(c["questions"]) for c in engine.calls] == [["q"], ["q"], ["q"]]
|
|
assert [list(c["questions"]["q"]["criteria"]) for c in engine.calls] == [
|
|
["o0", "o1", "o2"], ["o1", "o2", "o0"], ["o2", "o0", "o1"]]
|
|
assert out["id"] == "r1" and out["option_ids"] == ["o0", "o1", "o2"]
|
|
c = out["combined"]
|
|
assert (c["method"], c["orderings"]) == ("rotations", 3)
|
|
assert c["probabilities"] == pytest.approx([1 / 3] * 3)
|
|
assert c["top"] == "o0" # first maximum in the caller's order
|
|
assert c["agreement"] == pytest.approx(1 / 3)
|
|
assert [r["id"] for r in out["orderings"]] == ["r1#o0", "r1#o1", "r1#o2"]
|
|
assert [r["option_ids"] for r in out["orderings"]] == [["o0", "o1", "o2"], ["o1", "o2", "o0"], ["o2", "o0", "o1"]]
|
|
high, low = softmax_pair = (math.exp(2) / (math.exp(2) + 2), 1 / (math.exp(2) + 2))
|
|
assert c["spread"] == {i: [pytest.approx(low), pytest.approx(high)] for i in ("o0", "o1", "o2")}
|
|
del softmax_pair
|
|
|
|
|
|
def test_all_asks_every_permutation_and_is_422_above_4_options():
|
|
engine = FakeEngine()
|
|
client = make_client(engine)
|
|
out = client.post("/decide", json={**ROW, "options": opts(3), "orderings": "all"}, headers=AUTH).json()
|
|
assert len(engine.calls) == 6 and out["combined"]["orderings"] == 6
|
|
assert len({tuple(c["questions"]["q"]["criteria"]) for c in engine.calls}) == 6
|
|
assert list(engine.calls[0]["questions"]["q"]["criteria"]) == ["o0", "o1", "o2"]
|
|
response = client.post("/decide", json={**ROW, "options": opts(5), "orderings": "all"}, headers=AUTH)
|
|
assert response.status_code == 422 and len(engine.calls) == 6
|
|
|
|
|
|
def test_a_mixed_shared_request_runs_in_waves_and_wave_0_is_the_request_as_written():
|
|
engine = FakeEngine()
|
|
plain = {"state": "s", "decisions": [
|
|
{"id": "a", "question": "A?", "options": OPTIONS},
|
|
{"id": "b", "question": "B?", "options": opts(3)},
|
|
{"id": "c", "question": "C?", "options": OPTIONS}]}
|
|
client = make_client(engine)
|
|
plain_out = client.post("/decide/shared", json=plain, headers=AUTH).json()
|
|
plain_call = engine.calls.pop()
|
|
mixed = {**plain, "decisions": [plain["decisions"][0], {**plain["decisions"][1], "orderings": "rotations"},
|
|
plain["decisions"][2]]}
|
|
out = client.post("/decide/shared", json=mixed, headers=AUTH).json()
|
|
assert engine.calls[0] == plain_call # wave 0
|
|
assert [list(c["questions"]) for c in engine.calls] == [["q1", "q2", "q3"], ["q"], ["q"]]
|
|
assert [c["questions"]["q"]["instructions"] for c in engine.calls[1:]] == ["B?", "B?"]
|
|
assert out["results"][0] == plain_out["results"][0] and out["results"][2] == plain_out["results"][2]
|
|
assert out["results"][1]["combined"]["orderings"] == 3
|
|
assert out["timing"]["batch_size"] == 5 and out["timing"]["calls"] == 3
|
|
|
|
|
|
def test_waves_never_put_two_orderings_of_one_decision_in_one_call_and_are_packed_at_16():
|
|
engine = FakeEngine()
|
|
body = {"state": "s", "decisions": [{**d, "orderings": "rotations"} for d in decisions(17)]}
|
|
out = make_client(engine).post("/decide/shared", json=body, headers=AUTH).json()
|
|
assert [len(c["questions"]) for c in engine.calls] == [16, 1, 16, 1]
|
|
for call in engine.calls:
|
|
asked = [q["instructions"] for q in call["questions"].values()]
|
|
assert len(asked) == len(set(asked))
|
|
assert [r["id"] for r in out["results"]] == [f"decision-{i}" for i in range(17)]
|
|
assert out["timing"]["batch_size"] == 34
|
|
|
|
|
|
def test_orderings_count_toward_the_row_cap():
|
|
engine = FakeEngine()
|
|
body = {"state": "s", "decisions": [{**d, "orderings": "rotations"} for d in decisions(3)]}
|
|
assert make_client(engine, max_decisions=5).post("/decide/shared", json=body, headers=AUTH).status_code == 422
|
|
assert engine.calls == []
|
|
|
|
|
|
def test_a_zero_probability_is_floored_before_the_log():
|
|
class ZeroEngine(FakeEngine):
|
|
def predict(self, request):
|
|
response, sha = super().predict(request)
|
|
for answer in response["answers"].values():
|
|
first = next(iter(answer["probabilities"]))
|
|
answer["probabilities"] = {k: (0.0 if k == first else 1.0 / (len(answer["probabilities"]) - 1))
|
|
for k in answer["probabilities"]}
|
|
return response, sha
|
|
out = make_client(ZeroEngine()).post("/decide", json={**ROW, "orderings": "rotations"}, headers=AUTH)
|
|
assert out.status_code == 200
|
|
assert sum(out.json()["combined"]["probabilities"]) == pytest.approx(1.0)
|
|
|
|
|
|
import threading # noqa: E402
|
|
import time # noqa: E402
|
|
from concurrent.futures import ThreadPoolExecutor # noqa: E402
|
|
|
|
|
|
@pytest.mark.parametrize("chunked", [False, True])
|
|
def test_a_body_over_the_limit_is_413_whether_or_not_it_declares_its_length(chunked):
|
|
engine = FakeEngine()
|
|
client = make_client(engine, max_body_bytes=200)
|
|
body = ('{"id": "r1", "state": "' + "x" * 500 + '", "question": "Q?", "options": []}').encode()
|
|
content = (chunk for chunk in [body[:100], body[100:]]) if chunked else body
|
|
response = client.post("/decide", content=content, headers={**AUTH, "content-type": "application/json"})
|
|
assert response.status_code == 413
|
|
assert response.json()["error"]["code"] == "request_too_large"
|
|
assert engine.calls == []
|
|
|
|
|
|
def test_a_body_of_exactly_the_limit_is_accepted():
|
|
body = json.dumps(ROW).encode()
|
|
client = make_client(max_body_bytes=len(body))
|
|
assert client.post("/decide", content=body, headers={**AUTH, "content-type": "application/json"}).status_code == 200
|
|
|
|
|
|
def test_read_limited_stops_at_the_crossing_chunk_and_ignores_a_non_ascii_length():
|
|
import asyncio
|
|
from intern_decision_serve.app import ApiError, read_limited
|
|
consumed = []
|
|
|
|
async def chunks():
|
|
for i in range(10):
|
|
consumed.append(i)
|
|
yield b"x" * 100
|
|
|
|
with pytest.raises(ApiError) as info:
|
|
asyncio.run(read_limited(chunks(), None, 250))
|
|
assert info.value.status == 413 and consumed == [0, 1, 2]
|
|
|
|
async def small():
|
|
yield b"{}"
|
|
|
|
assert asyncio.run(read_limited(small(), "²", 100)) == b"{}" # int("²") would raise
|
|
|
|
|
|
class SlowEngine(FakeEngine):
|
|
"""Holds each call until released; records the peak number of calls inside at once and the
|
|
order in which requests' calls ran."""
|
|
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.inside = self.peak = 0
|
|
self.guard = threading.Lock()
|
|
self.release = threading.Event()
|
|
self.entered = threading.Event()
|
|
|
|
def predict(self, request):
|
|
with self.guard:
|
|
self.inside += 1
|
|
self.peak = max(self.peak, self.inside)
|
|
self.entered.set()
|
|
self.release.wait(5)
|
|
with self.guard:
|
|
self.inside -= 1
|
|
return super().predict(request)
|
|
|
|
|
|
def test_concurrent_requests_never_overlap_and_a_requests_chunks_are_not_interleaved():
|
|
engine = SlowEngine()
|
|
with make_client(engine) as client, ThreadPoolExecutor(4) as pool:
|
|
futures = [pool.submit(client.post, "/decide/shared",
|
|
json={"state": f"state {i}", "decisions": decisions(20)}, headers=AUTH)
|
|
for i in range(3)]
|
|
assert engine.entered.wait(5)
|
|
started = time.monotonic()
|
|
assert client.get("/health").status_code == 200
|
|
assert time.monotonic() - started < 1.0 # answered while a call is held
|
|
engine.release.set()
|
|
assert [f.result().status_code for f in futures] == [200] * 3
|
|
assert engine.peak == 1
|
|
states = [c["state"] for c in engine.calls]
|
|
assert all(states[k] == states[k + 1] for k in range(0, 6, 2)) # each request's 2 chunks back to back
|
|
|
|
|
|
def test_more_than_max_queue_requests_in_progress_get_429_busy_before_the_body_is_read():
|
|
engine = SlowEngine()
|
|
with make_client(engine, max_queue=2) as client, ThreadPoolExecutor(3) as pool:
|
|
held = [pool.submit(client.post, "/decide", json={**ROW, "id": f"r{i}"}, headers=AUTH) for i in range(2)]
|
|
assert engine.entered.wait(5)
|
|
time.sleep(0.2) # let the second request reach the queue
|
|
extra = client.post("/decide", content=b"{not even json", headers={**AUTH, "content-type": "application/json"})
|
|
assert extra.status_code == 429 and extra.json()["error"]["code"] == "busy"
|
|
engine.release.set()
|
|
assert [f.result().status_code for f in held] == [200, 200]
|
|
assert make_client(FakeEngine(), max_queue=2).post("/decide", json=ROW, headers=AUTH).status_code == 200
|