fix: take reported usage per count instead of merging usage dialects

This commit is contained in:
Ashen
2026-10-01 11:21:53 +05:30
parent 84bcce03b3
commit d8e8dd088c
6 changed files with 74 additions and 14 deletions
+1 -1
View File
@@ -1296,7 +1296,7 @@ async def _adjust_payment_for_tokens(
cache_creation_source=cost.cache_creation_source,
)
if terminal_usage is not None:
recorded_usage = {**(recorded_usage or {}), **terminal_usage}
recorded_usage = terminal_usage
presence = UsageFieldPresence(
input_source=cost.input_source,
output_source=cost.output_source,
+13 -7
View File
@@ -6,7 +6,7 @@ import secrets
import time
from dataclasses import dataclass
from ..payment.usage import NormalizedUsage, normalize_usage
from ..payment.usage import NormalizedUsage, normalize_usage, usage_field_presence
from .logging import get_logger
from .terminal_outcome_writer import (
TerminalOutcomeWriter,
@@ -55,16 +55,22 @@ def record_terminal_outcome(
) -> None:
"""Submit a settled outcome without awaiting storage or raising.
``usage`` is the raw upstream usage. It replaces the token counts and is
parsed here, so a malformed value can only mark a gap.
``usage`` is the raw upstream usage. Each count it reports, zero included,
replaces the matching argument and the rest are kept. It is parsed here, so
a malformed value can only mark a gap.
"""
try:
if usage is not None:
counted = normalize_usage(usage) or NormalizedUsage()
input_tokens = counted.input_tokens
output_tokens = counted.output_tokens
cache_read_input_tokens = counted.cache_read_tokens
cache_creation_input_tokens = counted.cache_write_tokens
reported = usage_field_presence(usage)
if reported.input_source != "missing":
input_tokens = counted.input_tokens
if reported.output_source != "missing":
output_tokens = counted.output_tokens
if reported.cache_read_source != "missing":
cache_read_input_tokens = counted.cache_read_tokens
if reported.cache_creation_source != "missing":
cache_creation_input_tokens = counted.cache_write_tokens
sources = {
name + "_source": getattr(context, name + "_source") or "missing"
for name in ("input", "output", "cache_read", "cache_creation")
@@ -141,8 +141,15 @@ def record_x_cashu_terminal_outcome(
metadata[name] = value
context = replace(context, **metadata)
if usage is not None:
# Stats keep what upstream reported, even where billing did not parse it.
presence = usage_field_presence(usage)
# Stats keep what upstream reported, even where billing did not parse it,
# and billing's labels stay on the counts upstream left out.
billed = UsageFieldPresence(
input_source=context.input_source or "missing",
output_source=context.output_source or "missing",
cache_read_source=context.cache_read_source or "missing",
cache_creation_source=context.cache_creation_source or "missing",
)
presence = billed.merged(usage_field_presence(usage))
context = replace(context, **presence.sources_dict())
counted: CostMetadata = cost_data if cost_data is not None else {}
record_terminal_outcome(
@@ -1383,9 +1383,10 @@ async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() ->
@pytest.mark.asyncio
@pytest.mark.parametrize("reported_output", [50, 0])
@pytest.mark.parametrize("api", ["chat", "responses"])
async def test_completed_stream_keeps_reported_usage_after_transport_error(
api: str,
api: str, reported_output: int
) -> None:
engine = await _engine()
async with AsyncSession(engine, expire_on_commit=False) as session:
@@ -1399,13 +1400,13 @@ async def test_completed_stream_keeps_reported_usage_after_transport_error(
[
b'data: {"model":"m","choices":[{"delta":{"content":"hi"},'
b'"finish_reason":"stop"}],"usage":{"prompt_tokens":100,'
b'"completion_tokens":50}}\n\n'
b'"completion_tokens":%d}}\n\n' % reported_output
]
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',
b'"usage":{"input_tokens":100,"output_tokens":%d}}}\n\n' % reported_output,
]
)
@@ -1483,7 +1484,7 @@ async def test_completed_stream_keeps_reported_usage_after_transport_error(
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 (counted.input_tokens, counted.output_tokens) == (100, reported_output)
assert (context.input_source, context.output_source) == ("reported", "reported")
await engine.dispose()
+25
View File
@@ -530,6 +530,31 @@ def test_record_wrapper_never_raises_on_invalid_or_failed_submission(
assert writer.submissions[0].model_identifier is None
def test_reported_usage_replaces_only_the_counts_it_reports(
monkeypatch: pytest.MonkeyPatch,
) -> None:
writer = MagicMock()
monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer)
for usage in (
{"prompt_tokens": 100, "completion_tokens": 0},
{"input_tokens": 100},
):
record_terminal_outcome(
TerminalOutcomeContext("request-c", "author/model"),
input_tokens=3,
output_tokens=2,
cache_read_input_tokens=0,
cache_creation_input_tokens=0,
revenue_msats=70,
usage=usage,
)
zero_output, input_only = (call.args[0] for call in writer.submit.call_args_list)
# A reported zero stands; a count upstream left out keeps billing's figure.
assert (zero_output.input_tokens, zero_output.output_tokens) == (100, 0)
assert (input_only.input_tokens, input_only.output_tokens) == (100, 2)
def test_outcome_rows_do_not_carry_the_request_id(
monkeypatch: pytest.MonkeyPatch,
) -> None:
+21
View File
@@ -62,6 +62,27 @@ def test_zero_usage_x_cashu_preserves_captured_presence() -> None:
assert record.call_args.kwargs["revenue_msats"] == 10_000
def test_x_cashu_outcome_keeps_billing_labels_for_counts_upstream_left_out() -> None:
writer = MagicMock()
with patch("routstr.core.terminal_outcomes.terminal_outcome_writer", writer):
record_x_cashu_terminal_outcome(
TerminalOutcomeContext("partial-usage", "author/model"),
{
"input_tokens": 3,
"output_tokens": 2,
"input_source": "estimated",
"output_source": "estimated",
},
amount=100,
unit="msat",
usage={"input_tokens": 100},
)
queued = writer.submit.call_args.args[0]
assert (queued.input_tokens, queued.input_source) == (100, "reported")
assert (queued.output_tokens, queued.output_source) == (2, "estimated")
# ---------------------------------------------------------------------------
# Non-streaming (chat completions)
# ---------------------------------------------------------------------------