feat(semif): 0.1.3 — order averaging, fast kernels, bug-hunt hardening (Prime)
Order averaging (Prime, after the 739aa03 spike):
- A decision may set orderings: rotations|all (all only for <= 4 options). Every
ordering goes to the engine in one shared batch.
- The reply keeps each native result and adds combined {probabilities (log-mean),
top, agreement, spread}.
- Through the service on SemIf's labelled sets (252 rows): 78.6% -> 88.1%
(group-bootstrap 95% CI +5.1..+14.3). Unanimous agreement is 94.5% accurate.
Fast kernels: flash-linear-attention 0.5.2 and causal-conv1d 1.7.0 are now the
default build. A/B on the empty GPU 3:
- parity with upstream went from 142/144 to 144/144;
- a ~2k-token /decide went from 169 to 92 ms server-side;
- short 3-rotation batches cost ~3-6 ms more.
triton builds a C shim at runtime, so the image carries gcc. Without it the
warm-up failed and startup failed closed.
Heid bug-hunt panel (4/4 arms, thread 01M3H3F4RR7XBP90KQ3A39H4SX), folded:
- Startup validation: VRAM cap 0 no longer means uncapped (C1); limits must be
>= 1 (S1); the token must be visible ASCII (S2); the calibration file must
exist and parse, with T in [0.05, 20] (S8, and C3's NaN leg).
- The body limit is checked before a chunk is kept, and a Unicode-digit
Content-Length no longer crashes (C2, S3).
- Failures while building the response now get the 500 envelope (C3).
- 429 busy past SEMIF_MAX_QUEUE requests in progress (C6).
- The engine releases memory on every non-validation failure, unchained after
gc; an empty OOM message is handled; 'out of memory' RuntimeErrors map to 503
(C4, C5, S9).
- The entry point forces HF_HUB_OFFLINE (S10). README wording fixed (S5, S6).
- New guard tests close the gaps the arms' mutation grids exposed: early stop of
the body read, a shared-route lock, calibration pass-through, the gc cycle,
the exact caps, TorchEngine.load's arch and device checks, and the offline
entry point.
86 tests.
Deployed on fv-ml1 GPU 1: parity 144/144, OOM and burst release verified, shared
capacity 63/51/26/16 rows at ~140/520/1960/3900 prefix tokens.
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
"""semif-serve HTTP behaviour against a fake engine (no torch, no model).
|
||||
Contract: services/semif-serve/semif-serve.contract.md"""
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
@@ -231,3 +233,104 @@ def test_health_reports_the_pins_limits_and_workloads():
|
||||
body = client.get("/health").json()
|
||||
assert body == {"status": "ok", "semif_commit": SEMIF_COMMIT, "model": FakeEngine().health(),
|
||||
"vram_cap_gib": 12.0, "max_tokens": 4096, "max_decisions": 8, "workloads": ["alerts", "triage"]}
|
||||
|
||||
|
||||
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_a_non_ascii_digit_content_length_is_ignored_not_a_crash():
|
||||
"""HTTP clients cannot send one (httpx refuses; h11 rejects it), so check the reader directly."""
|
||||
import asyncio
|
||||
from semif_serve.app import read_limited
|
||||
|
||||
async def body():
|
||||
yield b"{}"
|
||||
|
||||
assert asyncio.run(read_limited(body(), "²", 100)) == b"{}" # int("²") would raise
|
||||
|
||||
|
||||
def test_exactly_max_decisions_is_accepted():
|
||||
decisions = [{"id": str(i), "question": "Q?", "options": OPTIONS} for i in range(3)]
|
||||
client = make_client(max_decisions=3)
|
||||
assert client.post("/decide/shared", json={"state": "s", "decisions": decisions}, headers=AUTH).status_code == 200
|
||||
|
||||
|
||||
def test_calibration_leaves_every_native_field_alone():
|
||||
engine = FakeEngine(logits=(3.0, 1.0))
|
||||
body = make_client(engine, calibration={"triage": 2.0}).post(
|
||||
"/decide", json={**ROW, "workload": "triage"}, headers=AUTH).json()
|
||||
body.pop("calibrated")
|
||||
assert body == engine.direct(ROW)
|
||||
|
||||
|
||||
class MalformedEngine(FakeEngine):
|
||||
def direct(self, row):
|
||||
return {"id": row["id"], "option_ids": ["yes", "no"], "probabilities": [0.5, 0.5]} # no option_logits
|
||||
|
||||
def shared(self, rows):
|
||||
return [self.direct(r) for r in rows], {}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path, body", [
|
||||
("/decide", {**ROW, "workload": "triage"}),
|
||||
("/decide", {**ROW, "orderings": "rotations"}),
|
||||
])
|
||||
def test_a_malformed_scorer_result_is_an_envelope_500_not_a_bare_one(path, body):
|
||||
client = make_client(MalformedEngine(), calibration={"triage": 2.0})
|
||||
response = client.post(path, json=body, headers=AUTH)
|
||||
assert response.status_code == 500
|
||||
assert response.json()["error"]["code"] == "scoring_failed"
|
||||
|
||||
|
||||
class SharedSlowEngine(SlowEngine):
|
||||
def shared(self, rows):
|
||||
return [self.direct(r) for r in rows], {}
|
||||
|
||||
|
||||
def test_shared_requests_are_serialised_too():
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
engine = SharedSlowEngine()
|
||||
body = {"state": "s", "decisions": [{"id": "a", "question": "Q?", "options": OPTIONS}]}
|
||||
with make_client(engine) as client, ThreadPoolExecutor(3) as pool:
|
||||
futures = [pool.submit(client.post, "/decide/shared", json=body, headers=AUTH) for _ in range(3)]
|
||||
assert engine.entered.wait(5)
|
||||
engine.release.set()
|
||||
assert [f.result().status_code for f in futures] == [200] * 3
|
||||
assert engine.peak == 1
|
||||
|
||||
|
||||
def test_more_than_max_queue_requests_in_progress_get_429_busy():
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
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)
|
||||
import time
|
||||
deadline = time.monotonic() + 5
|
||||
while engine.inside + 0 < 1 and time.monotonic() < deadline:
|
||||
time.sleep(0.01)
|
||||
time.sleep(0.2) # let the second request reach the queue
|
||||
extra = client.post("/decide", json={**ROW, "id": "extra"}, headers=AUTH)
|
||||
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
|
||||
|
||||
|
||||
def test_read_limited_stops_reading_at_the_crossing_chunk():
|
||||
import asyncio
|
||||
from semif_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
|
||||
assert consumed == [0, 1, 2] # the third chunk crosses 250 and is never kept; nothing after is read
|
||||
|
||||
@@ -34,3 +34,45 @@ def test_a_calibration_table_needs_positive_finite_numbers(tmp_path, table):
|
||||
cal.write_text(json.dumps(table))
|
||||
with pytest.raises(ValueError, match="SEMIF_CALIBRATION"):
|
||||
Settings.from_env({"SEMIF_API_TOKEN": TOKEN, "SEMIF_CALIBRATION": str(cal)})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("var, value", [
|
||||
("SEMIF_VRAM_CAP_GIB", "0"), ("SEMIF_VRAM_CAP_GIB", "-4"), ("SEMIF_VRAM_CAP_GIB", "nan"), ("SEMIF_VRAM_CAP_GIB", "inf"),
|
||||
("SEMIF_MAX_TOKENS", "0"), ("SEMIF_MAX_DECISIONS", "-1"), ("SEMIF_MAX_BODY_BYTES", "0"), ("SEMIF_MAX_QUEUE", "0"),
|
||||
("SEMIF_MAX_TOKENS", "lots"), ("SEMIF_VRAM_CAP_GIB", "twelve"),
|
||||
])
|
||||
def test_out_of_range_or_unparseable_values_are_refused_naming_the_variable(var, value):
|
||||
with pytest.raises(ValueError, match=var):
|
||||
Settings.from_env({"SEMIF_API_TOKEN": TOKEN, var: value})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("token", ["x" * 31 + "\n", "x" * 30 + "\r\n", "x" * 32 + "\x00", "x" * 16 + " " + "x" * 16,
|
||||
"é" * 32])
|
||||
def test_a_token_with_non_visible_ascii_is_refused(token):
|
||||
with pytest.raises(ValueError, match="SEMIF_API_TOKEN"):
|
||||
Settings.from_env({"SEMIF_API_TOKEN": token})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("table", [{"w": True}, {"w": 1e-300}, {"w": 0.04}, {"w": 21}])
|
||||
def test_a_temperature_must_be_a_real_number_in_range(tmp_path, table):
|
||||
cal = tmp_path / "cal.json"
|
||||
cal.write_text(json.dumps(table))
|
||||
with pytest.raises(ValueError, match="SEMIF_CALIBRATION"):
|
||||
Settings.from_env({"SEMIF_API_TOKEN": TOKEN, "SEMIF_CALIBRATION": str(cal)})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("content", [None, "{not json"])
|
||||
def test_a_missing_or_malformed_calibration_file_is_a_named_startup_error(tmp_path, content):
|
||||
cal = tmp_path / "cal.json"
|
||||
if content is not None:
|
||||
cal.write_text(content)
|
||||
with pytest.raises(ValueError, match="SEMIF_CALIBRATION"):
|
||||
Settings.from_env({"SEMIF_API_TOKEN": TOKEN, "SEMIF_CALIBRATION": str(cal)})
|
||||
|
||||
|
||||
def test_in_range_edges_are_accepted(tmp_path):
|
||||
cal = tmp_path / "cal.json"
|
||||
cal.write_text(json.dumps({"lo": 0.05, "hi": 20}))
|
||||
s = Settings.from_env({"SEMIF_API_TOKEN": "!" + "~" * 31, "SEMIF_VRAM_CAP_GIB": "0.5", "SEMIF_MAX_QUEUE": "1",
|
||||
"SEMIF_CALIBRATION": str(cal)})
|
||||
assert (s.vram_cap_gib, s.max_queue, s.calibration) == (0.5, 1, {"lo": 0.05, "hi": 20.0})
|
||||
|
||||
@@ -7,7 +7,7 @@ import pytest
|
||||
|
||||
from semif_serve.config import Settings
|
||||
from semif_serve.engine import TorchEngine
|
||||
from semif_serve.errors import OutOfMemory
|
||||
from semif_serve.errors import OutOfMemory, ScoringFailed
|
||||
|
||||
|
||||
class FakeTorch:
|
||||
@@ -81,3 +81,33 @@ def test_a_burst_is_returned_to_the_driver_after_the_call(reserved_after, releas
|
||||
direct_fn=scorer, shared_fn=scorer, release_above_bytes=8 * 2**30 + 512 * 2**20)
|
||||
assert engine.direct({}) == {"ok": True}
|
||||
assert ReservingTorch.cuda.emptied == released
|
||||
|
||||
|
||||
def cyclic_tensor():
|
||||
"""A tensor held in a reference cycle, as real frames and tensors often are: only gc frees it."""
|
||||
t = Tensor()
|
||||
t.self_ref = t
|
||||
FakeTorch.cuda.watched.append(weakref.ref(t))
|
||||
return t
|
||||
|
||||
|
||||
@pytest.mark.parametrize("raised, expected_type, expected_message", [
|
||||
(lambda: FakeTorch.cuda.OutOfMemoryError(""), OutOfMemory, "CUDA out of memory"),
|
||||
(lambda: RuntimeError("CUBLAS_STATUS_ALLOC_FAILED: CUDA error: out of memory"), OutOfMemory,
|
||||
"CUBLAS_STATUS_ALLOC_FAILED: CUDA error: out of memory"),
|
||||
(lambda: RuntimeError("Invalid native prefix cache"), ScoringFailed, "RuntimeError: Invalid native prefix cache"),
|
||||
(lambda: KeyError("option_logits"), ScoringFailed, "KeyError: 'option_logits'"),
|
||||
])
|
||||
def test_every_non_validation_failure_is_released_unchained_after_gc(raised, expected_type, expected_message):
|
||||
FakeTorch.cuda.empties.clear(), FakeTorch.cuda.watched.clear()
|
||||
|
||||
def scorer(*_args):
|
||||
kv_cache = cyclic_tensor() # noqa: F841 — alive in this frame when it raises
|
||||
raise raised()
|
||||
|
||||
engine = TorchEngine(FakeTorch, None, None, {}, Settings(api_token="t" * 40), direct_fn=scorer, shared_fn=scorer)
|
||||
with pytest.raises(expected_type) as info:
|
||||
engine.shared([])
|
||||
assert str(info.value) == expected_message
|
||||
assert info.value.__cause__ is None and info.value.__context__ is None
|
||||
assert FakeTorch.cuda.empties == [True] # gc freed the cycle BEFORE the cache was emptied
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
"""TorchEngine.load and the entry point, against fake torch/semif modules (INV-3, INV-4, INV-5).
|
||||
The bug hunt found load() had no test at all: removing the arch check survived every test."""
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
from semif_serve.config import Settings
|
||||
|
||||
TOKEN = "t" * 40
|
||||
|
||||
|
||||
class Param:
|
||||
def __init__(self, device):
|
||||
self.device = types.SimpleNamespace(type=device)
|
||||
|
||||
|
||||
class Model:
|
||||
def __init__(self, device):
|
||||
self._device = device
|
||||
|
||||
def parameters(self):
|
||||
yield Param(self._device)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fakes(monkeypatch):
|
||||
calls = []
|
||||
cuda = types.SimpleNamespace(
|
||||
is_available=lambda: True,
|
||||
get_device_capability=lambda _i=0: (12, 0),
|
||||
get_arch_list=lambda: ["sm_90", "sm_120"],
|
||||
get_device_properties=lambda _i=0: types.SimpleNamespace(total_memory=96 * 2**30),
|
||||
set_per_process_memory_fraction=lambda f, _i=0: calls.append(("cap", round(f, 4))),
|
||||
memory_reserved=lambda _i=0: 8 * 2**30,
|
||||
OutOfMemoryError=type("OutOfMemoryError", (RuntimeError,), {}),
|
||||
empty_cache=lambda: calls.append(("empty",)),
|
||||
)
|
||||
torch = types.SimpleNamespace(cuda=cuda, __version__="2.10.0+cu128")
|
||||
state = {"device": "cuda", "warmup_raises": None}
|
||||
|
||||
def load_causal_model(model, revision, device, dtype):
|
||||
calls.append(("load", model, revision, device, dtype))
|
||||
return Model(state["device"]), object(), {"source": model}
|
||||
|
||||
def score(model, tok, row, meta, max_tokens):
|
||||
calls.append(("score", row["id"]))
|
||||
if state["warmup_raises"]:
|
||||
raise state["warmup_raises"]
|
||||
return {"id": row["id"]}
|
||||
|
||||
core = types.ModuleType("semif_phase1.core")
|
||||
core.load_causal_model = load_causal_model
|
||||
direct = types.ModuleType("semif_phase1.direct")
|
||||
direct.score = score
|
||||
shared = types.ModuleType("semif_phase1.shared")
|
||||
shared.score_shared = lambda *a: ([], {})
|
||||
pkg = types.ModuleType("semif_phase1")
|
||||
for name, mod in {"torch": torch, "semif_phase1": pkg, "semif_phase1.core": core,
|
||||
"semif_phase1.direct": direct, "semif_phase1.shared": shared}.items():
|
||||
monkeypatch.setitem(sys.modules, name, mod)
|
||||
return torch, calls, state
|
||||
|
||||
|
||||
def test_load_caps_before_the_weights_land_then_warms_up(fakes):
|
||||
from semif_serve.engine import TorchEngine
|
||||
_torch, calls, _ = fakes
|
||||
TorchEngine.load(Settings(api_token=TOKEN, vram_cap_gib=12.0))
|
||||
assert [c[0] for c in calls] == ["cap", "load", "score"]
|
||||
assert calls[0] == ("cap", round(12 / 96, 4))
|
||||
assert calls[2] == ("score", "semif-serve-warmup")
|
||||
|
||||
|
||||
def test_load_refuses_a_card_torch_has_no_kernels_for(fakes):
|
||||
from semif_serve.engine import TorchEngine
|
||||
torch, calls, _ = fakes
|
||||
torch.cuda.get_arch_list = lambda: ["sm_80", "sm_90"]
|
||||
with pytest.raises(RuntimeError, match="sm_120"):
|
||||
TorchEngine.load(Settings(api_token=TOKEN))
|
||||
assert not any(c[0] == "load" for c in calls)
|
||||
|
||||
|
||||
def test_load_refuses_a_model_that_landed_on_the_wrong_device(fakes):
|
||||
from semif_serve.engine import TorchEngine
|
||||
_torch, _calls, state = fakes
|
||||
state["device"] = "cpu"
|
||||
with pytest.raises(RuntimeError, match="landed on cpu"):
|
||||
TorchEngine.load(Settings(api_token=TOKEN))
|
||||
|
||||
|
||||
def test_load_fails_closed_when_the_warmup_decision_fails(fakes):
|
||||
from semif_serve.engine import TorchEngine
|
||||
from semif_serve.errors import ScoringFailed
|
||||
_torch, _calls, state = fakes
|
||||
state["warmup_raises"] = RuntimeError("Failed to find C compiler")
|
||||
with pytest.raises(ScoringFailed, match="C compiler"):
|
||||
TorchEngine.load(Settings(api_token=TOKEN))
|
||||
|
||||
|
||||
def test_the_entry_point_forces_offline_mode_before_the_engine_loads(fakes, monkeypatch):
|
||||
import semif_serve.engine as engine_mod
|
||||
from semif_serve import main
|
||||
seen = {}
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising=False)
|
||||
monkeypatch.setenv("SEMIF_API_TOKEN", TOKEN)
|
||||
monkeypatch.setattr(engine_mod.TorchEngine, "load",
|
||||
classmethod(lambda cls, s: seen.update(offline=os.environ.get("HF_HUB_OFFLINE")) or object()))
|
||||
main.app_from_env()
|
||||
assert seen["offline"] == "1"
|
||||
@@ -0,0 +1,132 @@
|
||||
"""Order averaging (0.1.3). Contract: semif-serve.contract.md § Order averaging."""
|
||||
import math
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from semif_serve.app import create_app
|
||||
from semif_serve.config import Settings
|
||||
|
||||
TOKEN = "t" * 40
|
||||
AUTH = {"Authorization": f"Bearer {TOKEN}"}
|
||||
OPTS = [{"id": "casual", "description": "A casual outing"},
|
||||
{"id": "date", "description": "A romantic date"},
|
||||
{"id": "booty", "description": "A booty call"}]
|
||||
ROW = {"id": "q", "state": "It's 2AM and I'm bored.", "question": "What is this?", "options": OPTS}
|
||||
BASE = {"casual": 1.0, "date": -1.0, "booty": 0.5} # what the model "really" thinks
|
||||
BIAS = 2.5 # added to whichever option is listed first
|
||||
|
||||
|
||||
def softmax(xs):
|
||||
top = max(xs)
|
||||
w = [math.exp(x - top) for x in xs]
|
||||
return [v / sum(w) for v in w]
|
||||
|
||||
|
||||
class BiasedEngine:
|
||||
"""Scores each option as BASE[id], plus BIAS for the first-listed option, the way a small model leans."""
|
||||
|
||||
def __init__(self):
|
||||
self.shared_calls, self.direct_calls = [], []
|
||||
|
||||
def health(self):
|
||||
return {}
|
||||
|
||||
def _score(self, row):
|
||||
ids = [o["id"] for o in row["options"]]
|
||||
logits = [BASE[i] + (BIAS if k == 0 else 0.0) for k, i in enumerate(ids)]
|
||||
return {"id": row["id"], "option_ids": ids, "probabilities": softmax(logits), "option_logits": logits,
|
||||
"prompt_sha256": row["id"], "probability_status": "uncalibrated"}
|
||||
|
||||
def direct(self, row):
|
||||
self.direct_calls.append(row)
|
||||
return self._score(row)
|
||||
|
||||
def shared(self, rows):
|
||||
self.shared_calls.append(rows)
|
||||
return [self._score(r) for r in rows], {"batch_size": len(rows)}
|
||||
|
||||
|
||||
def client(engine, **kw):
|
||||
return TestClient(create_app(Settings(api_token=TOKEN, **kw), engine))
|
||||
|
||||
|
||||
def test_rotations_send_one_shared_call_with_every_option_once_per_position_and_cancel_the_bias():
|
||||
engine = BiasedEngine()
|
||||
body = client(engine).post("/decide", json={**ROW, "orderings": "rotations"}, headers=AUTH).json()
|
||||
assert engine.direct_calls == [] and len(engine.shared_calls) == 1
|
||||
rows = engine.shared_calls[0]
|
||||
assert [r["id"] for r in rows] == ["q#o0", "q#o1", "q#o2"]
|
||||
assert [[o["id"] for o in r["options"]] for r in rows] == [
|
||||
["casual", "date", "booty"], ["date", "booty", "casual"], ["booty", "casual", "date"]]
|
||||
assert all(r["state"] == ROW["state"] and r["question"] == ROW["question"] for r in rows)
|
||||
|
||||
assert body["id"] == "q" and body["option_ids"] == ["casual", "date", "booty"]
|
||||
combined = body["combined"]
|
||||
assert combined["method"] == "rotations" and combined["orderings"] == 3
|
||||
# the first-position bias is additive, and every option sat first exactly once, so it cancels exactly
|
||||
assert combined["probabilities"] == pytest.approx(softmax([BASE["casual"], BASE["date"], BASE["booty"]]))
|
||||
assert combined["top"] == "casual"
|
||||
# the bias is strong enough that every ordering's first option wins: casual, date, booty
|
||||
tops = [max(zip(r["option_logits"], r["option_ids"]))[1] for r in body["orderings"]]
|
||||
assert tops == ["casual", "date", "booty"]
|
||||
assert combined["agreement"] == pytest.approx(1 / 3)
|
||||
assert body["orderings"] == [engine._score(r) for r in rows] # native results, unchanged
|
||||
for oid in ("casual", "date", "booty"):
|
||||
ps = [dict(zip(r["option_ids"], r["probabilities"]))[oid] for r in body["orderings"]]
|
||||
assert combined["spread"][oid] == pytest.approx([min(ps), max(ps)])
|
||||
|
||||
|
||||
def test_all_sends_every_permutation_with_the_callers_order_first():
|
||||
engine = BiasedEngine()
|
||||
body = client(engine).post("/decide", json={**ROW, "orderings": "all"}, headers=AUTH).json()
|
||||
rows = engine.shared_calls[0]
|
||||
orders = [tuple(o["id"] for o in r["options"]) for r in rows]
|
||||
assert len(orders) == 6 and len(set(orders)) == 6 and orders[0] == ("casual", "date", "booty")
|
||||
assert body["combined"]["orderings"] == 6 and body["combined"]["method"] == "all"
|
||||
|
||||
|
||||
def test_all_above_four_options_is_422_before_the_engine_runs():
|
||||
engine = BiasedEngine()
|
||||
five = [{"id": f"o{i}", "description": f"Option {i}"} for i in range(5)]
|
||||
r = client(engine).post("/decide", json={**ROW, "options": five, "orderings": "all"}, headers=AUTH)
|
||||
assert r.status_code == 422 and "rotations" in r.json()["error"]["message"]
|
||||
assert engine.shared_calls == []
|
||||
|
||||
|
||||
def test_a_mixed_shared_request_is_one_engine_call_with_results_in_request_order():
|
||||
engine = BiasedEngine()
|
||||
body = {"state": ROW["state"], "decisions": [
|
||||
{"id": "plain", "question": "Q1?", "options": OPTS},
|
||||
{"id": "avg", "question": "Q2?", "options": OPTS, "orderings": "rotations"},
|
||||
{"id": "plain2", "question": "Q3?", "options": OPTS[:2]}]}
|
||||
out = client(engine).post("/decide/shared", json=body, headers=AUTH).json()
|
||||
assert len(engine.shared_calls) == 1
|
||||
assert [r["id"] for r in engine.shared_calls[0]] == ["plain", "avg#o0", "avg#o1", "avg#o2", "plain2"]
|
||||
ids = [r["id"] for r in out["results"]]
|
||||
assert ids == ["plain", "avg", "plain2"]
|
||||
assert out["results"][0] == engine._score(engine.shared_calls[0][0]) # plain results keep the old shape
|
||||
assert "combined" in out["results"][1] and "combined" not in out["results"][2]
|
||||
assert out["timing"] == {"batch_size": 5}
|
||||
|
||||
|
||||
def test_expanded_rows_count_toward_the_cap():
|
||||
engine = BiasedEngine()
|
||||
r = client(engine, max_decisions=5).post("/decide/shared", json={"state": "s", "decisions": [
|
||||
{"id": "a", "question": "Q?", "options": OPTS, "orderings": "rotations"},
|
||||
{"id": "b", "question": "Q?", "options": OPTS, "orderings": "rotations"}]}, headers=AUTH)
|
||||
assert r.status_code == 422 and "6 scored rows" in r.json()["error"]["message"]
|
||||
assert engine.shared_calls == []
|
||||
|
||||
|
||||
def test_workload_with_orderings_is_422_and_workload_alone_still_calibrates():
|
||||
engine = BiasedEngine()
|
||||
c = client(engine, calibration={"triage": 2.0})
|
||||
r = c.post("/decide", json={**ROW, "orderings": "rotations", "workload": "triage"}, headers=AUTH)
|
||||
assert r.status_code == 422 and engine.shared_calls == []
|
||||
assert "calibrated" in c.post("/decide", json={**ROW, "workload": "triage"}, headers=AUTH).json()
|
||||
|
||||
|
||||
def test_an_unknown_orderings_value_is_422():
|
||||
r = client(BiasedEngine()).post("/decide", json={**ROW, "orderings": "shuffle"}, headers=AUTH)
|
||||
assert r.status_code == 422
|
||||
Reference in New Issue
Block a user