Reserve provider requests inside the durable worker envelope
Some checks failed
Governed runtime contract / contract (push) Failing after 31s
Some checks failed
Governed runtime contract / contract (push) Failing after 31s
Assistant: codex Assistant-Model: gpt-5.6-luna Assistant-Session: 01a07ff8-19d0-7820-b4d0-1353833cb7fc
This commit is contained in:
parent
4ae245a88f
commit
c30806b968
13 changed files with 1114 additions and 12 deletions
313
tests/test_request_admission.py
Normal file
313
tests/test_request_admission.py
Normal file
|
|
@ -0,0 +1,313 @@
|
|||
"""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"
|
||||
Loading…
Add table
Add a link
Reference in a new issue