diff --git a/routstr/upstream/terminal_outcome_tracking.py b/routstr/upstream/terminal_outcome_tracking.py index 55eaf5d8..69daba20 100644 --- a/routstr/upstream/terminal_outcome_tracking.py +++ b/routstr/upstream/terminal_outcome_tracking.py @@ -18,6 +18,15 @@ if TYPE_CHECKING: from ..payment.models import Model +def _event_usage(event: dict[str, Any], payload: dict[str, Any]) -> object: + usage = payload.get("usage") + if event.get("type") == "message_start" and isinstance(usage, dict): + # message_start's output count is a placeholder; message_delta carries + # the real one, so a stream stopped before it has no reported output. + return {key: value for key, value in usage.items() if key != "output_tokens"} + return usage + + @dataclass class TerminalOutcomeState: usage: dict[str, Any] | None = None @@ -26,8 +35,10 @@ class TerminalOutcomeState: # 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"), event.get("response"), event): - if isinstance(payload, dict) and isinstance(payload.get("usage"), dict): - self.usage = {**(self.usage or {}), **payload["usage"]} + if isinstance(payload, dict): + usage = _event_usage(event, payload) + if isinstance(usage, dict): + self.usage = {**(self.usage or {}), **usage} def terminal_outcome_context( @@ -51,7 +62,9 @@ def event_usage_presence(event: object) -> UsageFieldPresence: for key in ("message", "response"): nested = event.get(key) if isinstance(nested, dict): - presence = presence.merged(usage_field_presence(nested.get("usage"))) + presence = presence.merged( + usage_field_presence(_event_usage(event, nested)) + ) return presence diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 8a0e82f5..140183f2 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -1304,6 +1304,83 @@ async def test_native_messages_stats_ignore_network_chunk_boundaries( await engine.dispose() +@pytest.mark.asyncio +async def test_native_messages_hang_up_does_not_report_start_output() -> None: + engine = await _engine() + + @asynccontextmanager + async def sessions() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + async def chunks() -> AsyncGenerator[bytes, None]: + # message_start carries a placeholder output count; the real one only + # arrives in message_delta, which this client never waits for. + yield ( + b'event: message_start\ndata: {"type":"message_start","message":' + b'{"model":"test-model","usage":{"input_tokens":10,"output_tokens":1}}}\n\n' + ) + yield ( + b'event: message_delta\ndata: {"type":"message_delta","delta":' + b'{"stop_reason":"end_turn"},"usage":{"output_tokens":5}}\n\n' + ) + + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + upstream_response.aiter_bytes = chunks + upstream_response.aclose = AsyncMock() + writer = MagicMock() + provider = BaseUpstreamProvider("https://unused.example", "test-key") + try: + async with sessions() as session: + key = ApiKey(hashed_key="messages-hang-up", balance=1_000) + session.add(key) + await session.commit() + await pay_for_request(key, 100, session) + reservation = await get_reservation_snapshot(key, session) + + with ( + patch("routstr.upstream.base.create_session", sessions), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + auth_module.adjust_payment_for_tokens, + ), + patch("routstr.core.terminal_outcomes.terminal_outcome_writer", writer), + patch( + "routstr.payment.cost_calculation._get_pricing_rates", + return_value=(1_000.0, 1_000.0, 1_000.0, 1_000.0, "configured"), + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005 + ), + patch("routstr.auth.ROUTSTR_FEE_PERCENT", 0), + patch("routstr.upstream.count_tokens._count_with_litellm", return_value=3), + patch( + "routstr.upstream.count_tokens._count_text_with_litellm", return_value=0 + ), + ): + response = await provider.handle_streaming_messages_completion( + upstream_response, + key, + 100, + reservation_snapshot=reservation, + terminal_outcome=TerminalOutcomeContext( + "messages-hang-up", "test-model" + ), + ) + body = cast(AsyncGenerator[bytes, None], response.body_iterator) + await anext(body) + await body.aclose() + + writer.submit.assert_called_once() + outcome = writer.submit.call_args.args[0] + assert (outcome.input_tokens, outcome.input_source) == (10, "reported") + assert outcome.output_source == "estimated" + finally: + await engine.dispose() + + def test_stream_cut_inside_a_character_does_not_raise() -> None: state = TerminalOutcomeState() tail = 'data: {"delta":{"text":"日本'.encode()[:-1]