fix: record every billed request, not only completed ones

This commit is contained in:
Ashen
2026-10-01 14:40:05 +05:30
parent d8e8dd088c
commit 8bc575125c
3 changed files with 61 additions and 64 deletions
@@ -77,12 +77,6 @@ class TerminalOutcomeState:
def settlement_context(
self, *, require_success: bool = False
) -> TerminalOutcomeContext | None:
if self.failure_seen:
return None
if self.transport_failed and not self.success_marker_seen:
return None
if require_success and not self.success_marker_seen:
return None
return self.context
@@ -30,7 +30,6 @@ from routstr.upstream.base import BaseUpstreamProvider
from routstr.upstream.terminal_outcome_tracking import (
TerminalOutcomeState,
observe_terminal_sse_bytes,
track_generic_terminal_stream,
)
@@ -41,25 +40,6 @@ async def _engine() -> AsyncEngine:
return engine
@pytest.mark.asyncio
async def test_generic_terminal_outcome_requires_consumed_clean_eof() -> None:
context = TerminalOutcomeContext(
outcome_id="generic-terminal",
model_identifier="test-model",
)
state = TerminalOutcomeState(context)
assert state.settlement_context(require_success=True) is None
async def chunks() -> AsyncGenerator[bytes, None]:
yield b"complete"
assert [
chunk async for chunk in track_generic_terminal_stream(chunks(), state)
] == [b"complete"]
assert state.settlement_context(require_success=True) is context
@pytest.mark.asyncio
async def test_release_reservation_is_durable_and_idempotent() -> None:
engine = await _engine()
@@ -364,6 +344,47 @@ async def test_generic_stream_abort_settles_and_closes_once() -> None:
response.aclose.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_aborted_generic_stream_still_records_its_settlement() -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
session = MagicMock()
session.get = AsyncMock(return_value=MagicMock(spec=ApiKey))
session_context = MagicMock()
session_context.__aenter__ = AsyncMock(return_value=session)
session_context.__aexit__ = AsyncMock(return_value=None)
adjust = AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0})
terminal_outcome = TerminalOutcomeContext(
outcome_id="generic-aborted", model_identifier="test-model"
)
upstream_response = _opaque_stream_response(b"first", b"second")
upstream_response.status_code = 200
upstream_response.headers = {}
with (
patch("routstr.upstream.base.adjust_payment_for_tokens", adjust),
patch("routstr.upstream.base.create_session", return_value=session_context),
):
response = provider._generic_streaming_response(
upstream_response,
"key-hash",
500,
"audio/speech",
None,
provider.provider_fee,
MagicMock(spec=ReservationSnapshot),
terminal_outcome,
)
stream = cast(AsyncGenerator[bytes, None], response.body_iterator)
assert await anext(stream) == b"first"
await stream.aclose()
adjust.assert_awaited_once()
assert adjust.await_args is not None
assert adjust.await_args.kwargs["terminal_outcome"] is terminal_outcome
@pytest.mark.asyncio
async def test_streaming_response_closes_iterator_when_downstream_send_is_cancelled() -> (
None
@@ -677,23 +698,16 @@ async def test_responses_streaming_releases_and_raises_on_billing_failure(
@pytest.mark.asyncio
@pytest.mark.parametrize("api", ["chat", "responses"])
@pytest.mark.parametrize("finalization_fails", [False, True])
@pytest.mark.parametrize("terminal_marker_seen", [False, True])
async def test_partial_remote_protocol_error_finalizes_and_closes_once(
api: str,
finalization_fails: bool,
terminal_marker_seen: bool,
) -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
if terminal_marker_seen and api == "chat":
yield b'data: {"model":"test","choices":[{"finish_reason":"stop"}]}\n\n'
elif terminal_marker_seen:
yield b'data: {"type":"response.completed","response":{"status":"completed"}}\n\n'
else:
yield b'data: {"model":"test","choices":[{"delta":{"content":"hi"}}]}\n\n'
yield b'data: {"model":"test","choices":[{"delta":{"content":"hi"}}]}\n\n'
raise httpx.RemoteProtocolError("incomplete chunked read")
upstream_response = MagicMock(
@@ -754,9 +768,7 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
adjust.assert_awaited_once()
assert adjust.await_args is not None
assert adjust.await_args.kwargs["terminal_outcome"] is (
terminal_outcome if terminal_marker_seen else None
)
assert adjust.await_args.kwargs["terminal_outcome"] is terminal_outcome
if finalization_fails:
session.rollback.assert_awaited_once()
release.assert_awaited_once_with(snapshot, session, 500)
@@ -768,7 +780,7 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
@pytest.mark.asyncio
@pytest.mark.parametrize("api", ["chat", "responses", "messages"])
async def test_stream_closed_before_first_chunk_is_not_a_completed_outcome(
async def test_stream_closed_before_first_chunk_still_records_its_settlement(
api: str,
) -> None:
provider = BaseUpstreamProvider(
@@ -814,7 +826,7 @@ async def test_stream_closed_before_first_chunk_is_not_a_completed_outcome(
adjust.assert_awaited_once()
assert adjust.await_args is not None
assert adjust.await_args.kwargs["terminal_outcome"] is None
assert adjust.await_args.kwargs["terminal_outcome"] is terminal_outcome
upstream_response.aclose.assert_awaited_once()
@@ -1122,7 +1134,7 @@ async def test_gemini_messages_finalizes_when_response_start_fails() -> None:
@pytest.mark.asyncio
async def test_native_messages_split_error_event_is_not_counted() -> None:
async def test_native_messages_error_event_still_records_its_settlement() -> None:
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key"
)
@@ -1169,7 +1181,7 @@ async def test_native_messages_split_error_event_is_not_counted() -> None:
adjust.assert_awaited_once()
assert adjust.await_args is not None
assert adjust.await_args.kwargs["terminal_outcome"] is None
assert adjust.await_args.kwargs["terminal_outcome"] is terminal_outcome
@pytest.mark.asyncio
@@ -1327,7 +1339,6 @@ def test_stream_cut_inside_a_character_does_not_raise() -> None:
tail = 'data: {"delta":{"text":"日本'.encode()[:-1]
assert observe_terminal_sse_bytes(state, b"", tail, final=True) == b""
assert state.settlement_context() is None
def test_routstr_upstream_cost_event_is_not_provider_usage() -> None:
@@ -1490,10 +1501,9 @@ async def test_completed_stream_keeps_reported_usage_after_transport_error(
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal_marker_seen", [False, True])
async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat(
terminal_marker_seen: bool,
) -> None:
async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() -> (
None
):
"""A client abort releases the hold after charging only estimated usage.
Starlette closes the response generator (``aclose``) on disconnect. The
@@ -1517,10 +1527,7 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat(
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
# A live stream that never sends a usage chunk or [DONE]; the client
# disconnects after the first delta.
finish_reason = b',"finish_reason":"stop"' if terminal_marker_seen else b""
yield (
b'data: {"choices":[{"delta":{"content":"hi"}' + finish_reason + b"}]}\n\n"
)
yield b'data: {"choices":[{"delta":{"content":"hi"}}]}\n\n'
yield b'data: {"choices":[{"delta":{"content":" there"}}]}\n\n'
upstream_response = MagicMock(
@@ -1600,19 +1607,15 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat(
# 3 input tokens × 10 msats + 2 output tokens × 20 msats = 70 msats.
assert final_key.total_spent == 70
assert final_key.balance == 930
if terminal_marker_seen:
record_outcome.assert_called_once()
recorded_context = record_outcome.call_args.args[0]
assert recorded_context.outcome_id == terminal_outcome.outcome_id
assert recorded_context.model_identifier == terminal_outcome.model_identifier
assert recorded_context.input_source == "estimated"
assert recorded_context.output_source == "estimated"
assert recorded_context.cache_read_source == "missing"
assert recorded_context.cache_creation_source == "missing"
assert record_outcome.call_args.kwargs["revenue_msats"] == 70
else:
# Billing settles, but an early disconnect is not a completed outcome.
record_outcome.assert_not_called()
record_outcome.assert_called_once()
recorded_context = record_outcome.call_args.args[0]
assert recorded_context.outcome_id == terminal_outcome.outcome_id
assert recorded_context.model_identifier == terminal_outcome.model_identifier
assert recorded_context.input_source == "estimated"
assert recorded_context.output_source == "estimated"
assert recorded_context.cache_read_source == "missing"
assert recorded_context.cache_creation_source == "missing"
assert record_outcome.call_args.kwargs["revenue_msats"] == 70
# The heartbeat is gone — no forever-renewing task on an abandoned request.
assert snapshot.release_id not in auth_module._reservation_heartbeats
await engine.dispose()
+1 -1
View File
@@ -434,4 +434,4 @@ async def test_truncated_json_tail_on_connection_close() -> None:
# non-JSON ever reached the client.
assert contents == ["ok"]
assert adjustment.await_args is not None
assert adjustment.await_args.kwargs["terminal_outcome"] is None
assert adjustment.await_args.kwargs["terminal_outcome"] is terminal_outcome