refactor: drop stream finish tracking the ledger no longer reads

This commit is contained in:
Ashen
2026-10-01 14:40:14 +05:30
parent 8bc575125c
commit d42a1cff75
3 changed files with 88 additions and 271 deletions
+81 -161
View File
@@ -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),
+2 -67
View File
@@ -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)
@@ -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'