mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix: keep reported usage when a completed stream loses its connection
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user