diff --git a/routstr/upstream/terminal_outcome_tracking.py b/routstr/upstream/terminal_outcome_tracking.py index a98091dd..f35497bf 100644 --- a/routstr/upstream/terminal_outcome_tracking.py +++ b/routstr/upstream/terminal_outcome_tracking.py @@ -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 diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index acdc4aad..892189c1 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -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() diff --git a/tests/unit/test_streaming_sse_providers.py b/tests/unit/test_streaming_sse_providers.py index 37d090e5..fbbfe747 100644 --- a/tests/unit/test_streaming_sse_providers.py +++ b/tests/unit/test_streaming_sse_providers.py @@ -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