mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: keep message_start's placeholder output count out of reported stats
This commit is contained in:
@@ -18,6 +18,15 @@ if TYPE_CHECKING:
|
||||
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
|
||||
class TerminalOutcomeState:
|
||||
usage: dict[str, Any] | None = None
|
||||
@@ -26,8 +35,10 @@ class TerminalOutcomeState:
|
||||
# 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"), event.get("response"), event):
|
||||
if isinstance(payload, dict) and isinstance(payload.get("usage"), dict):
|
||||
self.usage = {**(self.usage or {}), **payload["usage"]}
|
||||
if isinstance(payload, dict):
|
||||
usage = _event_usage(event, payload)
|
||||
if isinstance(usage, dict):
|
||||
self.usage = {**(self.usage or {}), **usage}
|
||||
|
||||
|
||||
def terminal_outcome_context(
|
||||
@@ -51,7 +62,9 @@ def event_usage_presence(event: object) -> UsageFieldPresence:
|
||||
for key in ("message", "response"):
|
||||
nested = event.get(key)
|
||||
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
|
||||
|
||||
|
||||
|
||||
@@ -1304,6 +1304,83 @@ async def test_native_messages_stats_ignore_network_chunk_boundaries(
|
||||
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:
|
||||
state = TerminalOutcomeState()
|
||||
tail = 'data: {"delta":{"text":"日本'.encode()[:-1]
|
||||
|
||||
Reference in New Issue
Block a user