"""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"