fix: keep message_start's placeholder output count out of reported stats

This commit is contained in:
Ashen
2026-10-02 01:59:46 +05:30
parent b0f0ca808e
commit 23c1cef7d8
2 changed files with 93 additions and 3 deletions
+16 -3
View File
@@ -18,6 +18,15 @@ if TYPE_CHECKING:
from ..payment.models import Model 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 @dataclass
class TerminalOutcomeState: class TerminalOutcomeState:
usage: dict[str, Any] | None = None usage: dict[str, Any] | None = None
@@ -26,8 +35,10 @@ class TerminalOutcomeState:
# Messages report input and output usage in separate events, and a # Messages report input and output usage in separate events, and a
# completed Responses event nests its usage under "response". # completed Responses event nests its usage under "response".
for payload in (event.get("message"), event.get("response"), event): for payload in (event.get("message"), event.get("response"), event):
if isinstance(payload, dict) and isinstance(payload.get("usage"), dict): if isinstance(payload, dict):
self.usage = {**(self.usage or {}), **payload["usage"]} usage = _event_usage(event, payload)
if isinstance(usage, dict):
self.usage = {**(self.usage or {}), **usage}
def terminal_outcome_context( def terminal_outcome_context(
@@ -51,7 +62,9 @@ def event_usage_presence(event: object) -> UsageFieldPresence:
for key in ("message", "response"): for key in ("message", "response"):
nested = event.get(key) nested = event.get(key)
if isinstance(nested, dict): 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 return presence
@@ -1304,6 +1304,83 @@ async def test_native_messages_stats_ignore_network_chunk_boundaries(
await engine.dispose() 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: def test_stream_cut_inside_a_character_does_not_raise() -> None:
state = TerminalOutcomeState() state = TerminalOutcomeState()
tail = 'data: {"delta":{"text":"日本'.encode()[:-1] tail = 'data: {"delta":{"text":"日本'.encode()[:-1]