From 84bcce03b3982d85d4497e6ca5a9595b2192a734 Mon Sep 17 00:00:00 2001 From: Ashen <310210685+ashen0x@users.noreply.github.com> Date: Sun, 27 Sep 2026 20:49:36 +0530 Subject: [PATCH] fix: keep reported usage when a completed stream loses its connection --- routstr/upstream/base.py | 2 + routstr/upstream/terminal_outcome_tracking.py | 5 +- .../test_streaming_billing_finalization.py | 107 ++++++++++++++++++ 3 files changed, 112 insertions(+), 2 deletions(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index efb45b63..3d10a690 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1222,6 +1222,7 @@ class BaseUpstreamProvider: terminal_outcome=outcome_state.settlement_context( require_success=True ), + terminal_usage=outcome_state.usage, ) usage_finalized = True except Exception: @@ -1723,6 +1724,7 @@ class BaseUpstreamProvider: terminal_outcome=outcome_state.settlement_context( require_success=True ), + terminal_usage=outcome_state.usage, ) usage_finalized = True except Exception: diff --git a/routstr/upstream/terminal_outcome_tracking.py b/routstr/upstream/terminal_outcome_tracking.py index f790ca43..9809edc2 100644 --- a/routstr/upstream/terminal_outcome_tracking.py +++ b/routstr/upstream/terminal_outcome_tracking.py @@ -33,8 +33,9 @@ class TerminalOutcomeState: nested_response = event.get("response") if isinstance(nested_response, dict): status = str(nested_response.get("status") or status).lower() - # Messages report input and output usage in separate events. - for payload in (event.get("message"), event): + # Messages report input and output usage in separate events, and a + # completed Responses event nests its usage under "response". + for payload in (event.get("message"), nested_response, event): if isinstance(payload, dict) and isinstance(payload.get("usage"), dict): self.usage = {**(self.usage or {}), **payload["usage"]} if ( diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index c7b4126d..863e3834 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -25,6 +25,7 @@ from routstr.core.db import ApiKey, ReservationRelease from routstr.core.terminal_outcomes import TerminalOutcomeContext from routstr.payment.cost_calculation import MaxCostData from routstr.payment.models import Architecture, Model, Pricing +from routstr.payment.usage import normalize_usage from routstr.upstream.base import BaseUpstreamProvider from routstr.upstream.terminal_outcome_tracking import ( TerminalOutcomeState, @@ -1381,6 +1382,112 @@ async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> await engine.dispose() +@pytest.mark.asyncio +@pytest.mark.parametrize("api", ["chat", "responses"]) +async def test_completed_stream_keeps_reported_usage_after_transport_error( + api: str, +) -> None: + engine = await _engine() + async with AsyncSession(engine, expire_on_commit=False) as session: + key = ApiKey(hashed_key=f"{api}-reported-then-cut", balance=10_000) + session.add(key) + await session.commit() + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) + + events = ( + [ + b'data: {"model":"m","choices":[{"delta":{"content":"hi"},' + b'"finish_reason":"stop"}],"usage":{"prompt_tokens":100,' + b'"completion_tokens":50}}\n\n' + ] + if api == "chat" + else [ + b'data: {"type":"response.output_text.delta","delta":"hi"}\n\n', + b'data: {"type":"response.completed","response":{"status":"completed",' + b'"usage":{"input_tokens":100,"output_tokens":50}}}\n\n', + ] + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + for event in events: + yield event + raise httpx.RemoteProtocolError("incomplete chunked read") + + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + upstream_response.aiter_bytes = aiter_bytes + upstream_response.aclose = AsyncMock() + model = Model( + id="test-model", + name="test-model", + created=0, + description="", + context_length=8_192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing(prompt=0.01, completion=0.02), + sats_pricing=Pricing(prompt=0.01, completion=0.02), + ) + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key", provider_fee=1.0 + ) + record_outcome = MagicMock() + try: + with ( + patch( + "routstr.upstream.base.create_session", + side_effect=lambda: AsyncSession(engine, expire_on_commit=False), + ), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + auth_module.adjust_payment_for_tokens, + ), + patch("routstr.auth.record_terminal_outcome", record_outcome), + patch("routstr.upstream.count_tokens._count_with_litellm", return_value=3), + patch( + "routstr.upstream.count_tokens._count_text_with_litellm", + return_value=2, + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", + return_value=5.0e-5, + ), + ): + handler = getattr(provider, f"handle_streaming_{api}_completion") + response = await handler( + response=upstream_response, + key=key, + max_cost_for_model=500, + model_obj=model, + reservation_snapshot=snapshot, + request_body=json.dumps({"model": model.id, "messages": []}).encode(), + terminal_outcome=TerminalOutcomeContext(f"{api}-outcome", model.id), + ) + async for _ in response.body_iterator: + pass + finally: + await auth_module._stop_reservation_heartbeat(snapshot.release_id) + + async with AsyncSession(engine, expire_on_commit=False) as session: + charged = await session.get(ApiKey, key.hashed_key) + # Billing keeps its fallback estimate: 3 input x 10 + 2 output x 20 msats. + assert charged is not None and charged.total_spent == 70 + record_outcome.assert_called_once() + context = record_outcome.call_args.args[0] + counted = normalize_usage(record_outcome.call_args.kwargs["usage"]) + assert counted is not None + assert (counted.input_tokens, counted.output_tokens) == (100, 50) + assert (context.input_source, context.output_source) == ("reported", "reported") + await engine.dispose() + + @pytest.mark.asyncio @pytest.mark.parametrize("terminal_marker_seen", [False, True]) async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat(