diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 3d10a690..6b03f0f0 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -100,7 +100,6 @@ from .terminal_outcome_tracking import ( observe_terminal_sse_bytes, record_x_cashu_terminal_outcome, terminal_outcome_context, - track_generic_terminal_stream, ) if typing.TYPE_CHECKING: @@ -1199,7 +1198,7 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None provider_seen: str | None = None - outcome_state = TerminalOutcomeState(terminal_outcome) + outcome_state = TerminalOutcomeState() async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1219,9 +1218,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context( - require_success=True - ), + terminal_outcome=terminal_outcome, terminal_usage=outcome_state.usage, ) usage_finalized = True @@ -1312,7 +1309,6 @@ class BaseUpstreamProvider: if data.strip() == b"[DONE]": done_seen = True - outcome_state.mark_success() return obj = json_codec.loads(data) @@ -1370,7 +1366,6 @@ class BaseUpstreamProvider: # mid-event, so ``data`` is incomplete JSON. Emitting it # as a ``data:`` frame would hand the client invalid # JSON (the "unexpected token" parse error). Drop it. - outcome_state.mark_transport_failure() return # Non-JSON data payload (partial fragment already reassembled # by buffering, or a provider control string). Re-prefix each @@ -1414,7 +1409,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context(), + terminal_outcome=terminal_outcome, ) usage_finalized = True except BaseException as e: @@ -1478,7 +1473,6 @@ class BaseUpstreamProvider: yield b"data: [DONE]\n\n" except httpx.RemoteProtocolError as stream_error: - outcome_state.mark_transport_failure() logger.warning( "Upstream stream ended before the response was complete", extra={ @@ -1486,8 +1480,7 @@ class BaseUpstreamProvider: "key_hash": key.hashed_key[:8] + "...", }, ) - except BaseException as stream_error: - outcome_state.mark_transport_failure() + except Exception as stream_error: logger.warning( "Streaming interrupted; finalizing before closing upstream", extra={ @@ -1549,8 +1542,6 @@ class BaseUpstreamProvider: try: content = await response.aread() response_json = json.loads(content) - outcome_state = TerminalOutcomeState(terminal_outcome) - outcome_state.observe(response_json) self._apply_provider_field(response_json) logger.debug( @@ -1584,7 +1575,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context(), + terminal_outcome=terminal_outcome, ) await session.refresh(key) @@ -1701,7 +1692,7 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None provider_seen: str | None = None - outcome_state = TerminalOutcomeState(terminal_outcome) + outcome_state = TerminalOutcomeState() async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1721,9 +1712,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context( - require_success=True - ), + terminal_outcome=terminal_outcome, terminal_usage=outcome_state.usage, ) usage_finalized = True @@ -1800,7 +1789,6 @@ class BaseUpstreamProvider: if data.strip() == b"[DONE]": done_seen = True - outcome_state.mark_success() return obj = json_codec.loads(data) @@ -1836,7 +1824,6 @@ class BaseUpstreamProvider: # Final flush of a truncated tail: upstream closed # mid-event, so ``data`` is incomplete JSON. Dropping it # avoids handing the client an invalid ``data:`` frame. - outcome_state.mark_transport_failure() return # Re-prefix each line so multi-line ``data`` stays valid SSE # framing for the client. @@ -1874,7 +1861,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context(), + terminal_outcome=terminal_outcome, ) usage_finalized = True except BaseException as e: @@ -1938,7 +1925,6 @@ class BaseUpstreamProvider: yield b"data: [DONE]\n\n" except httpx.RemoteProtocolError as stream_error: - outcome_state.mark_transport_failure() logger.warning( "Upstream Responses API stream ended before the response was complete", extra={ @@ -1946,8 +1932,7 @@ class BaseUpstreamProvider: "key_hash": key.hashed_key[:8] + "...", }, ) - except BaseException as stream_error: - outcome_state.mark_transport_failure() + except Exception as stream_error: logger.warning( "Responses API streaming interrupted; finalizing before closing upstream", extra={ @@ -2008,8 +1993,6 @@ class BaseUpstreamProvider: try: content = await response.aread() response_json = json.loads(content) - outcome_state = TerminalOutcomeState(terminal_outcome) - outcome_state.observe(response_json) self._apply_provider_field(response_json) logger.debug( @@ -2044,7 +2027,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context(), + terminal_outcome=terminal_outcome, ) await session.refresh(key) @@ -2134,7 +2117,7 @@ class BaseUpstreamProvider: model_obj: Model | None, provider_fee: float | None, reservation_snapshot: ReservationSnapshot, - outcome_state: TerminalOutcomeState | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> None: """Finalize payment for a generic streaming request.""" async with create_session() as session: @@ -2158,11 +2141,7 @@ class BaseUpstreamProvider: model_obj=model_obj, provider_fee=provider_fee, reservation_snapshot=reservation_snapshot, - terminal_outcome=outcome_state.settlement_context( - require_success=True - ) - if outcome_state is not None - else None, + terminal_outcome=terminal_outcome, ) logger.debug( "Finalized generic streaming payment", @@ -2191,7 +2170,6 @@ class BaseUpstreamProvider: provider_fee: float | None, reservation_snapshot: ReservationSnapshot, finalizer: PersistentStreamFinalizer | None = None, - outcome_state: TerminalOutcomeState | None = None, ) -> AsyncGenerator[bytes, None]: """Relay an opaque stream and settle it even if the caller disconnects.""" if finalizer is None: @@ -2204,16 +2182,12 @@ class BaseUpstreamProvider: model_obj, provider_fee, reservation_snapshot, - outcome_state, ), response, ) ) - chunks = response.aiter_bytes() - if outcome_state is not None: - chunks = track_generic_terminal_stream(chunks, outcome_state) try: - async for chunk in chunks: + async for chunk in response.aiter_bytes(): yield chunk finally: await finalizer.run() @@ -2229,7 +2203,6 @@ class BaseUpstreamProvider: reservation_snapshot: ReservationSnapshot, terminal_outcome: TerminalOutcomeContext | None = None, ) -> ClosingStreamingResponse: - outcome_state = TerminalOutcomeState(terminal_outcome) finalizer = PersistentStreamFinalizer( lambda: finalize_and_close_stream( lambda: self._finalize_generic_streaming_payment( @@ -2239,7 +2212,7 @@ class BaseUpstreamProvider: model_obj, provider_fee, reservation_snapshot, - outcome_state, + terminal_outcome, ), response, ) @@ -2253,7 +2226,6 @@ class BaseUpstreamProvider: provider_fee, reservation_snapshot, finalizer, - outcome_state, ) return ClosingStreamingResponse( stream, @@ -2278,11 +2250,9 @@ class BaseUpstreamProvider: last_model_seen: str | None = None provider_seen: str | None = None usage_presence = UsageFieldPresence() - outcome_state = TerminalOutcomeState(terminal_outcome) + outcome_state = TerminalOutcomeState() - async def finalize_without_usage( - *, require_success: bool = False - ) -> bytes | None: + async def finalize_without_usage() -> bytes | None: nonlocal usage_finalized if usage_finalized: return None @@ -2300,9 +2270,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context( - require_success=require_success - ), + terminal_outcome=terminal_outcome, usage_presence=usage_presence, terminal_usage=outcome_state.usage, ) @@ -2326,7 +2294,7 @@ class BaseUpstreamProvider: async def finalize_db_only() -> None: if not usage_finalized: - await finalize_without_usage(require_success=True) + await finalize_without_usage() stream_finalizer = PersistentStreamFinalizer( lambda: finalize_and_close_stream(finalize_db_only, response) @@ -2533,7 +2501,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context(), + terminal_outcome=terminal_outcome, usage_presence=usage_presence, terminal_usage=outcome_state.usage, ) @@ -2569,12 +2537,10 @@ class BaseUpstreamProvider: yield maybe_cost_event except httpx.ReadError: - outcome_state.mark_transport_failure() if not usage_finalized: await finalize_without_usage() # Upstream dropped the connection mid-stream; response already started, swallow silently - except BaseException: - outcome_state.mark_transport_failure() + except Exception: if not usage_finalized: await finalize_without_usage() raise @@ -2608,8 +2574,6 @@ class BaseUpstreamProvider: try: content = await response.aread() response_json = json.loads(content) - outcome_state = TerminalOutcomeState(terminal_outcome) - outcome_state.observe(response_json) if requested_model: if "model" in response_json: @@ -2639,7 +2603,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context(), + terminal_outcome=terminal_outcome, ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2761,8 +2725,6 @@ class BaseUpstreamProvider: ) response_json = messages_dispatch.coerce_litellm_payload(result) - outcome_state = TerminalOutcomeState(terminal_outcome) - outcome_state.observe(response_json) if requested_model and "model" in response_json: response_json["model"] = requested_model if not isinstance(response_json.get("usage"), dict): @@ -2780,7 +2742,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context(), + terminal_outcome=terminal_outcome, ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2829,10 +2791,7 @@ class BaseUpstreamProvider: ) response_json = messages_dispatch.coerce_litellm_payload(result) - outcome_state = TerminalOutcomeState( - terminal_outcome_context(request_id, model_obj) - ) - outcome_state.observe(response_json) + terminal_outcome = terminal_outcome_context(request_id, model_obj) self._apply_provider_field(response_json) if requested_model and "model" in response_json: response_json["model"] = requested_model @@ -2865,10 +2824,9 @@ class BaseUpstreamProvider: request_id=request_id, ) except BaseException: - if outcome_state.settlement_context() is not None: - mark_terminal_outcome_loss( - "X-Cashu LiteLLM refund commit ambiguous" - ) + mark_terminal_outcome_loss( + "X-Cashu LiteLLM refund commit ambiguous" + ) raise response_headers["X-Cashu"] = refund_token refund_amount_sent = refund_amount @@ -2881,15 +2839,13 @@ class BaseUpstreamProvider: }, ) - terminal_context = outcome_state.settlement_context() - if terminal_context is not None: - record_x_cashu_terminal_outcome( - terminal_context, - cost_data, - amount=amount, - unit=unit, - refund_amount=refund_amount_sent, - ) + record_x_cashu_terminal_outcome( + terminal_outcome, + cost_data, + amount=amount, + unit=unit, + refund_amount=refund_amount_sent, + ) return Response( content=json.dumps(response_json).encode(), @@ -2918,11 +2874,8 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None usage_presence = UsageFieldPresence() - outcome_state = TerminalOutcomeState(terminal_outcome) - async def finalize_without_usage( - *, require_success: bool = False - ) -> bytes | None: + async def finalize_without_usage() -> bytes | None: nonlocal usage_finalized if usage_finalized: return None @@ -2952,9 +2905,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context( - require_success=require_success - ), + terminal_outcome=terminal_outcome, usage_presence=usage_presence, ) usage_finalized = True @@ -2980,7 +2931,7 @@ class BaseUpstreamProvider: async def finalize_stream() -> None: try: if not usage_finalized: - await finalize_without_usage(require_success=True) + await finalize_without_usage() finally: await aclose_if_needed(iterator) @@ -3000,7 +2951,6 @@ class BaseUpstreamProvider: async for annotated in messages_dispatch.stream_annotated_events( iterator, requested_model ): - outcome_state.observe(annotated.event) usage_presence = usage_presence.merged( event_usage_presence(annotated.event) ) @@ -3065,7 +3015,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, - terminal_outcome=outcome_state.settlement_context(), + terminal_outcome=terminal_outcome, usage_presence=usage_presence, ) self.inject_cost_metadata( @@ -3099,8 +3049,7 @@ class BaseUpstreamProvider: if cost_event is not None: yield cost_event - except BaseException: - outcome_state.mark_transport_failure() + except Exception: if not usage_finalized: await finalize_without_usage() raise @@ -3148,14 +3097,11 @@ class BaseUpstreamProvider: total_cost = 0.0 input_cost = 0.0 output_cost = 0.0 - outcome_state = TerminalOutcomeState( - terminal_outcome_context(request_id, model_obj) - ) + terminal_outcome = terminal_outcome_context(request_id, model_obj) async for annotated in messages_dispatch.stream_annotated_events( iterator, requested_model ): - outcome_state.observe(annotated.event) usage_presence = usage_presence.merged( event_usage_presence(annotated.event) ) @@ -3253,17 +3199,13 @@ class BaseUpstreamProvider: }, ) except asyncio.CancelledError: - if outcome_state.settlement_context() is not None: - mark_terminal_outcome_loss( - "X-Cashu LiteLLM stream settlement cancelled" - ) + mark_terminal_outcome_loss( + "X-Cashu LiteLLM stream settlement cancelled" + ) raise except Exception as exc: settlement_failed = True - if outcome_state.settlement_context() is not None: - mark_terminal_outcome_loss( - "X-Cashu LiteLLM stream settlement failed" - ) + mark_terminal_outcome_loss("X-Cashu LiteLLM stream settlement failed") logger.error( "Error calculating cost for streaming /v1/messages", extra={ @@ -3275,15 +3217,13 @@ class BaseUpstreamProvider: ) if not settlement_failed: - terminal_context = outcome_state.settlement_context() - if terminal_context is not None: - record_x_cashu_terminal_outcome( - replace(terminal_context, **usage_presence.sources_dict()), - cost_data, - amount=amount, - unit=unit, - refund_amount=refund_amount_sent, - ) + record_x_cashu_terminal_outcome( + replace(terminal_outcome, **usage_presence.sources_dict()), + cost_data, + amount=amount, + unit=unit, + refund_amount=refund_amount_sent, + ) if cost_data: _inject_cost_response_headers(response_headers, cost_data) @@ -4286,7 +4226,7 @@ class BaseUpstreamProvider: usage_estimator = MissingUsageEstimator(request_body, model_obj) refund_amount_sent = 0 settlement_failed = False - outcome_state = TerminalOutcomeState(terminal_outcome) + outcome_state = TerminalOutcomeState() # Stats observe both SSE prefix forms; billing keeps its existing parse. observe_terminal_sse_bytes(outcome_state, b"", content_str.encode(), final=True) @@ -4408,12 +4348,12 @@ class BaseUpstreamProvider: # inputMsats/outputMsats/totalMsats for x-cashu requests. _inject_cost_response_headers(response_headers, cost_data) except asyncio.CancelledError: - if outcome_state.settlement_context() is not None: + if terminal_outcome is not None: mark_terminal_outcome_loss("X-Cashu streaming settlement cancelled") raise except Exception as e: settlement_failed = True - if outcome_state.settlement_context() is not None: + if terminal_outcome is not None: mark_terminal_outcome_loss("X-Cashu streaming settlement failed") logger.error( "Error calculating cost for streaming response", @@ -4427,10 +4367,9 @@ class BaseUpstreamProvider: ) if not settlement_failed: - terminal_context = outcome_state.settlement_context() - if terminal_context is not None: + if terminal_outcome is not None: record_x_cashu_terminal_outcome( - terminal_context, + terminal_outcome, cost_data, amount=amount, unit=unit, @@ -4501,8 +4440,6 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) - outcome_state = TerminalOutcomeState(terminal_outcome) - outcome_state.observe(response_json) self._apply_provider_field(response_json) _apply_estimated_usage( response_json, request_body, model_obj, amount, unit, "chat" @@ -4583,7 +4520,7 @@ class BaseUpstreamProvider: request_id=request_id, ) except BaseException: - if outcome_state.settlement_context() is not None: + if terminal_outcome is not None: mark_terminal_outcome_loss("X-Cashu refund commit ambiguous") raise response_headers["X-Cashu"] = refund_token @@ -4599,10 +4536,9 @@ class BaseUpstreamProvider: }, ) - terminal_context = outcome_state.settlement_context() - if terminal_context is not None: + if terminal_outcome is not None: record_x_cashu_terminal_outcome( - terminal_context, + terminal_outcome, cost_data, amount=amount, unit=unit, @@ -5409,13 +5345,10 @@ class BaseUpstreamProvider: usage_estimator = MissingUsageEstimator(request_body, model_obj) refund_amount_sent = 0 settlement_failed = False - outcome_state = TerminalOutcomeState( - terminal_outcome_context(request_id, model_obj) - ) + terminal_outcome = terminal_outcome_context(request_id, model_obj) for _fields, data in events: if data.strip() == "[DONE]": - outcome_state.mark_success() continue try: data_json = json.loads(data) @@ -5424,7 +5357,6 @@ class BaseUpstreamProvider: if not isinstance(data_json, dict): continue usage_estimator.observe(data_json) - outcome_state.observe(data_json) # Canonical Responses API events carry model and usage nested under # "response" (response.completed/incomplete); older shapes put them # at the top level. @@ -5532,15 +5464,11 @@ class BaseUpstreamProvider: # inputMsats/outputMsats/totalMsats for x-cashu requests. _inject_cost_response_headers(response_headers, cost_data) except asyncio.CancelledError: - if outcome_state.settlement_context() is not None: - mark_terminal_outcome_loss( - "X-Cashu Responses stream settlement cancelled" - ) + mark_terminal_outcome_loss("X-Cashu Responses stream settlement cancelled") raise except Exception as e: settlement_failed = True - if outcome_state.settlement_context() is not None: - mark_terminal_outcome_loss("X-Cashu Responses stream settlement failed") + mark_terminal_outcome_loss("X-Cashu Responses stream settlement failed") logger.error( "Error calculating cost for streaming Responses API response", extra={ @@ -5553,16 +5481,14 @@ class BaseUpstreamProvider: ) if not settlement_failed: - terminal_context = outcome_state.settlement_context() - if terminal_context is not None: - record_x_cashu_terminal_outcome( - terminal_context, - cost_data, - amount=amount, - unit=unit, - refund_amount=refund_amount_sent, - usage=usage_data, - ) + record_x_cashu_terminal_outcome( + terminal_outcome, + cost_data, + amount=amount, + unit=unit, + refund_amount=refund_amount_sent, + usage=usage_data, + ) provider_seen: str | None = None for i, (fields, data) in enumerate(events): @@ -5615,10 +5541,7 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) - outcome_state = TerminalOutcomeState( - terminal_outcome_context(request_id, model_obj) - ) - outcome_state.observe(response_json) + terminal_outcome = terminal_outcome_context(request_id, model_obj) self._apply_provider_field(response_json) _apply_estimated_usage( response_json, request_body, model_obj, amount, unit, "responses" @@ -5696,10 +5619,9 @@ class BaseUpstreamProvider: request_id=request_id, ) except BaseException: - if outcome_state.settlement_context() is not None: - mark_terminal_outcome_loss( - "X-Cashu Responses refund commit ambiguous" - ) + mark_terminal_outcome_loss( + "X-Cashu Responses refund commit ambiguous" + ) raise response_headers["X-Cashu"] = refund_token @@ -5714,16 +5636,14 @@ class BaseUpstreamProvider: }, ) - terminal_context = outcome_state.settlement_context() - if terminal_context is not None: - record_x_cashu_terminal_outcome( - terminal_context, - cost_data, - amount=amount, - unit=unit, - refund_amount=max(0, refund_amount), - usage=response_json.get("usage"), - ) + record_x_cashu_terminal_outcome( + terminal_outcome, + cost_data, + amount=amount, + unit=unit, + refund_amount=max(0, refund_amount), + usage=response_json.get("usage"), + ) return Response( content=json.dumps(response_json), diff --git a/routstr/upstream/terminal_outcome_tracking.py b/routstr/upstream/terminal_outcome_tracking.py index f35497bf..55eaf5d8 100644 --- a/routstr/upstream/terminal_outcome_tracking.py +++ b/routstr/upstream/terminal_outcome_tracking.py @@ -1,9 +1,8 @@ -"""Capture how upstream responses end, for the settled outcome ledger.""" +"""Capture upstream usage and outcome contexts for the settled outcome ledger.""" from __future__ import annotations import json -from collections.abc import AsyncGenerator, AsyncIterator from dataclasses import dataclass, replace from typing import TYPE_CHECKING, Any @@ -21,63 +20,14 @@ if TYPE_CHECKING: @dataclass class TerminalOutcomeState: - context: TerminalOutcomeContext | None - success_marker_seen: bool = False - failure_seen: bool = False - transport_failed: bool = False usage: dict[str, Any] | None = None def observe(self, event: dict[str, Any]) -> None: - event_type = str(event.get("type") or "").lower() - status = str(event.get("status") or "").lower() - 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, and a # completed Responses event nests its usage under "response". - for payload in (event.get("message"), nested_response, event): + 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 ( - event.get("error") is not None - or event_type in {"error", "response.failed"} - or status in {"cancelled", "failed"} - ): - self.failure_seen = True - return - choices = event.get("choices") - if isinstance(choices, list): - finish_reasons = { - str(choice.get("finish_reason") or "").lower() - for choice in choices - if isinstance(choice, dict) and choice.get("finish_reason") is not None - } - if "error" in finish_reasons: - self.failure_seen = True - return - if finish_reasons - {""}: - self.success_marker_seen = True - # An output-limit truncation is a paid terminal response, like length. - if event_type in { - "response.completed", - "response.incomplete", - "message_stop", - } or status in {"completed", "incomplete"}: - self.success_marker_seen = True - delta = event.get("delta") - if isinstance(delta, dict) and delta.get("stop_reason") not in (None, ""): - self.success_marker_seen = True - - def mark_success(self) -> None: - self.success_marker_seen = True - - def mark_transport_failure(self) -> None: - self.transport_failed = True - - def settlement_context( - self, *, require_success: bool = False - ) -> TerminalOutcomeContext | None: - return self.context def terminal_outcome_context( @@ -159,18 +109,6 @@ def record_x_cashu_terminal_outcome( ) -async def track_generic_terminal_stream( - stream: AsyncIterator[bytes], state: TerminalOutcomeState -) -> AsyncGenerator[bytes, None]: - try: - async for chunk in stream: - yield chunk - state.mark_success() - except BaseException: - state.mark_transport_failure() - raise - - def observe_terminal_sse_bytes( state: TerminalOutcomeState, buffered: bytes, @@ -206,14 +144,11 @@ def observe_terminal_sse_bytes( continue payload = b"\n".join(data_lines) if payload.strip() == b"[DONE]": - state.mark_success() continue try: parsed = json.loads(payload) except ValueError: # Bytes cut inside a character raise UnicodeDecodeError, not JSONDecodeError. - if final: - state.mark_transport_failure() continue if isinstance(parsed, dict): state.observe(parsed) diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 892189c1..89714d28 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -3,7 +3,7 @@ import json from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from typing import cast -from unittest.mock import ANY, AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest @@ -304,7 +304,6 @@ async def test_generic_stream_completion_settles_and_closes_once() -> None: None, provider.provider_fee, reservation, - None, ) response.aclose.assert_awaited_once_with() @@ -339,7 +338,6 @@ async def test_generic_stream_abort_settles_and_closes_once() -> None: None, provider.provider_fee, reservation, - None, ) response.aclose.assert_awaited_once_with() @@ -444,7 +442,7 @@ async def test_streaming_response_closes_iterator_when_downstream_send_is_cancel None, provider.provider_fee, reservation, - ANY, + None, ) upstream_response.aclose.assert_awaited_once_with() @@ -500,7 +498,7 @@ async def test_generic_stream_settles_when_response_start_fails() -> None: None, provider.provider_fee, reservation, - ANY, + None, ) upstream_response.aclose.assert_awaited_once_with() @@ -1300,51 +1298,15 @@ async def test_native_messages_stats_ignore_network_chunk_boundaries( await engine.dispose() -def test_anthropic_stop_reason_is_terminal_before_transport_failure() -> None: - terminal_outcome = TerminalOutcomeContext( - outcome_id="messages-stop-reason", - model_identifier="test-model", - ) - state = TerminalOutcomeState(terminal_outcome) - - state.observe( - { - "type": "message_delta", - "delta": {"stop_reason": "end_turn"}, - "usage": {"output_tokens": 2}, - } - ) - state.mark_transport_failure() - - assert state.settlement_context() is terminal_outcome - - -def test_responses_output_limit_is_terminal_before_transport_failure() -> None: - terminal_outcome = TerminalOutcomeContext( - outcome_id="responses-incomplete", - model_identifier="test-model", - ) - state = TerminalOutcomeState(terminal_outcome) - - state.observe({"type": "response.incomplete", "response": {"status": "incomplete"}}) - state.mark_transport_failure() - - assert state.settlement_context() is terminal_outcome - - def test_stream_cut_inside_a_character_does_not_raise() -> None: - state = TerminalOutcomeState( - TerminalOutcomeContext(outcome_id="cut", model_identifier="test-model") - ) + state = TerminalOutcomeState() tail = 'data: {"delta":{"text":"日本'.encode()[:-1] assert observe_terminal_sse_bytes(state, b"", tail, final=True) == b"" def test_routstr_upstream_cost_event_is_not_provider_usage() -> None: - state = TerminalOutcomeState( - TerminalOutcomeContext(outcome_id="routstr-upstream", model_identifier="m") - ) + state = TerminalOutcomeState() stream = ( b'event: message_start\ndata: {"type":"message_start","message":{"usage":' b'{"input_tokens":100,"cache_read_input_tokens":1000,"output_tokens":1}}}\n\n'