fix: keep reported usage when a completed stream loses its connection

This commit is contained in:
Ashen
2026-10-01 11:21:53 +05:30
parent ef429769eb
commit 84bcce03b3
3 changed files with 112 additions and 2 deletions
+2
View File
@@ -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:
@@ -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 (
@@ -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(