rein-aharness/tests/test_request_admission.py
tegwick c30806b968
Some checks failed
Governed runtime contract / contract (push) Failing after 31s
Reserve provider requests inside the durable worker envelope
Assistant: codex
Assistant-Model: gpt-5.6-luna
Assistant-Session: 01a07ff8-19d0-7820-b4d0-1353833cb7fc
2026-09-09 21:50:12 +02:00

313 lines
11 KiB
Python

"""Real SQLite/HTTP, deterministic provider; no paid calls or credentials."""
import http.client
import http.server
import json
import threading
from concurrent.futures import ThreadPoolExecutor
from datetime import UTC, datetime, timedelta
import pytest
from llm_connect.messages_gate import MessagesPolicy, MessagesServer
from rein_aharness.request_admission import RequestLedger
from rein_aharness.spend_admission import SpendAdmissionError, SpendLedger
from test_spend_admission import ledger as ledger # shared private ledger fixture
from test_spend_admission import run, success
@pytest.fixture
def meter(ledger):
child = RequestLedger(ledger)
child.initialize()
ledger.reserve(run())
return child
@pytest.fixture
def policy():
return MessagesPolicy(
"fixture:not-live-tariff", "fixture-model", 1000, 1000, 1000, 1000
)
def route(meter, policy, **kwargs):
return meter.bind_route(
run().id,
policy.sha256,
lease_id="fixture-lease",
expires_at=(datetime.now(UTC) + timedelta(seconds=60)).isoformat(),
**kwargs,
)
def test_parent_capacity_no_refund_replay_and_reopen(meter, policy):
token = route(meter, policy)
receipt = meter.reserve_request(token, policy.sha256, 2_000_000)
with pytest.raises(SpendAdmissionError, match="unresolved"):
RequestLedger(meter.parent).reserve_request(token, policy.sha256, 1)
meter.complete_request(receipt, 10)
with pytest.raises(SpendAdmissionError, match="replay"):
meter.complete_request(receipt, 10)
second = meter.reserve_request(token, policy.sha256, 2_000_000)
meter.complete_request(second, 0)
with pytest.raises(SpendAdmissionError, match="capacity"):
meter.reserve_request(token, policy.sha256, 2_000_000)
assert sum(row["liability_microusd"] for row in meter.status()) == 4_000_000
assert meter.parent.status()["reservations"][0]["liability"] == 5_000_000
serialized = json.dumps(meter.status())
assert token not in serialized and "private prompt" not in serialized
def test_concurrent_admission_only_one_pending(meter, policy):
token = route(meter, policy)
def reserve(_):
try:
RequestLedger(
SpendLedger(meter.parent.path, meter.parent.policy)
).reserve_request(token, policy.sha256, 1_000_000)
return True
except SpendAdmissionError:
return False
with ThreadPoolExecutor(max_workers=8) as pool:
assert sum(pool.map(reserve, range(8))) == 1
assert len(meter.status()) == 1
def test_unknown_holds_parent_even_after_success_and_operator_closes(meter, policy):
token = route(meter, policy)
meter.reserve_request(token, policy.sha256, 100)
assert not meter.parent.observe(run().id, success(meter.parent, run()))
with pytest.raises(SpendAdmissionError, match="revoked"):
meter.reserve_request(token, policy.sha256, 100)
assert meter.parent.status()["reservations"][0]["state"] == "held"
meter.parent.reconcile(
run().id, cost_usd="1", receipt="audit:provider-stopped-final"
)
assert meter.status()[0]["state"] == "charged"
with pytest.raises(SpendAdmissionError):
meter.bind_route(
run().id,
policy.sha256,
lease_id="replacement",
expires_at="2098-01-01T00:00:00Z",
)
def test_unknown_terminal_revokes_and_known_overrun_freezes(meter, policy):
token = route(meter, policy)
receipt = meter.reserve_request(token, policy.sha256, 100)
with pytest.raises(SpendAdmissionError, match="breached"):
meter.complete_request(receipt, 6_000_000)
assert meter.parent.status()["breached"]
assert meter.parent.status()["reservations"][0]["liability"] == 6_000_000
assert not meter.parent.observe(run().id, {"evidence": {}})
assert not meter.request_active(receipt)
with pytest.raises(SpendAdmissionError):
meter.reserve_request(token, policy.sha256, 1)
def test_wrong_token_policy_expired_lease_and_revocation(meter, policy):
now = datetime.now(UTC)
token = meter.bind_route(
run().id,
policy.sha256,
lease_id="fixture",
expires_at=(now + timedelta(seconds=1)).isoformat(),
now=now,
)
for credential, digest in [("a" * 43, policy.sha256), (token, "b" * 64)]:
with pytest.raises(SpendAdmissionError):
meter.reserve_request(credential, digest, 1, now=now)
receipt = meter.reserve_request(token, policy.sha256, 1, now=now)
assert not meter.request_active(receipt, now=now + timedelta(seconds=2))
meter.revoke_route(run().id)
assert not meter.request_active(receipt, now=now + timedelta(seconds=2))
def stream_bytes(*, truncated=False, output=10):
events = [
{
"type": "message_start",
"message": {
"model": "fixture-model",
"usage": {
"input_tokens": 10,
"output_tokens": 0,
"cache_creation_input_tokens": 20,
"cache_read_input_tokens": 30,
},
},
},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": output},
},
{"type": "message_stop"},
]
if truncated:
events.pop()
return "".join(
"event: " + e["type"] + "\ndata: " + json.dumps(e) + "\n\n" for e in events
).encode()
@pytest.fixture
def fake_provider():
seen = []
mode = {"status": 200, "truncated": False, "output": 10}
class Handler(http.server.BaseHTTPRequestHandler):
def log_message(self, *args):
pass
def do_POST(self):
body = self.rfile.read(int(self.headers["Content-Length"]))
seen.append(
{
"key": self.headers.get("x-api-key"),
"path": self.path,
"body": json.loads(body),
"authorization": self.headers.get("Authorization"),
}
)
response = stream_bytes(truncated=mode["truncated"], output=mode["output"])
self.send_response(mode["status"])
self.send_header("Content-Type", "text/event-stream")
self.send_header("Content-Length", str(len(response)))
self.end_headers()
if wait := mode.get("wait"):
mode["entered"].set()
wait.wait(5)
self.wfile.write(response)
server = http.server.ThreadingHTTPServer(("127.0.0.1", 0), Handler)
threading.Thread(target=server.serve_forever, daemon=True).start()
yield server, seen, mode
server.shutdown()
server.server_close()
@pytest.fixture
def gateway(meter, policy, fake_provider):
provider, seen, mode = fake_provider
token = route(meter, policy)
server = MessagesServer(
policy,
meter,
provider_key="dummy-owner-key",
upstream_url=f"http://127.0.0.1:{provider.server_port}",
allow_test_http=True,
)
server.start()
yield server, token, seen, mode
server.stop()
def request(gateway, **changes):
server, token, _, _ = gateway
body = {
"model": "fixture-model",
"max_tokens": 1000,
"stream": True,
"messages": [{"role": "user", "content": "private fixture prompt"}],
}
body.update(changes)
connection = http.client.HTTPConnection("127.0.0.1", server.port, timeout=5)
connection.request(
"POST",
"/v1/messages?beta=true",
json.dumps(body),
{"Content-Type": "application/json", "x-api-key": token},
)
response = connection.getresponse()
result = response.status, response.read()
connection.close()
return result
def test_transport_reserves_before_forward_and_preserves_parent(gateway, meter):
for _ in range(2):
assert request(gateway)[0] == 200
assert request(gateway)[0] == 400
assert len(gateway[2]) == 2
assert all(
r["key"] == "dummy-owner-key" and not r["authorization"] for r in gateway[2]
)
assert all(r["observed_microusd"] == 70000 for r in meter.status())
assert sum(r["liability_microusd"] for r in meter.status()) == 4_000_000
assert meter.parent.observe(run().id, success(meter.parent, run()))
assert request(gateway)[0] == 400 and len(gateway[2]) == 2
@pytest.mark.parametrize(
"changes",
[
{"max_tokens": True},
{"max_tokens": 1001},
{"max_tokens": 0},
{"max_tokens": 1.5},
{"model": "other-model"},
{"stream": False},
{"service_tier": "auto"},
{"tools": [{"type": "web_search_20250305", "name": "web_search"}]},
{
"messages": [
{
"role": "user",
"content": [
{
"type": "image",
"source": {"type": "url", "url": "https://example.com"},
}
],
}
]
},
{"context_management": {}},
{"cache_control": {"type": "unknown"}},
],
)
def test_unsupported_requests_never_forward(gateway, meter, changes):
assert request(gateway, **changes)[0] == 400
assert gateway[2] == [] and meter.status() == []
@pytest.mark.parametrize("failure", ["truncated", "429", "overrun"])
def test_uncertain_or_overrun_response_prevents_retry(gateway, meter, failure):
if failure == "truncated":
gateway[3]["truncated"] = True
if failure == "429":
gateway[3]["status"] = 429
if failure == "overrun":
gateway[3]["output"] = 6000
request(gateway)
assert request(gateway)[0] == 400
assert len(gateway[2]) == 1
assert meter.status()[0]["state"] == "held"
assert not meter.parent.observe(run().id, success(meter.parent, run()))
def test_inflight_lease_loss_holds_and_denies_concurrent_forward(gateway, meter):
gateway[3].update(wait=threading.Event(), entered=threading.Event())
with ThreadPoolExecutor() as pool:
first = pool.submit(request, gateway)
assert gateway[3]["entered"].wait(3)
assert request(gateway)[0] == 400
meter.revoke_route(run().id)
gateway[3]["wait"].set()
first.result(5)
assert len(gateway[2]) == 1 and meter.status()[0]["state"] == "held"
def test_missing_request_schema_never_recreates_parent(ledger, policy):
ledger.reserve(run())
meter = RequestLedger(ledger)
with pytest.raises(SpendAdmissionError):
route(meter, policy)
with pytest.raises(SpendAdmissionError, match="before dispatch"):
meter.initialize()
assert ledger.status()["reservations"][0]["state"] == "held"