mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix: record every billed request, not only completed ones
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user