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
|
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]
|
||||||
|
|||||||
Reference in New Issue
Block a user