From d8e8dd088c6326fd3a24344f6458249e3c136342 Mon Sep 17 00:00:00 2001 From: Ashen <310210685+ashen0x@users.noreply.github.com> Date: Sun, 27 Sep 2026 21:44:40 +0530 Subject: [PATCH] fix: take reported usage per count instead of merging usage dialects --- routstr/auth.py | 2 +- routstr/core/terminal_outcomes.py | 20 +++++++++------ routstr/upstream/terminal_outcome_tracking.py | 11 ++++++-- .../test_streaming_billing_finalization.py | 9 ++++--- tests/unit/test_terminal_outcomes.py | 25 +++++++++++++++++++ tests/unit/test_x_cashu_cost_sats.py | 21 ++++++++++++++++ 6 files changed, 74 insertions(+), 14 deletions(-) diff --git a/routstr/auth.py b/routstr/auth.py index 4ccab255..3381e65a 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -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, diff --git a/routstr/core/terminal_outcomes.py b/routstr/core/terminal_outcomes.py index 03c52cad..6103d981 100644 --- a/routstr/core/terminal_outcomes.py +++ b/routstr/core/terminal_outcomes.py @@ -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") diff --git a/routstr/upstream/terminal_outcome_tracking.py b/routstr/upstream/terminal_outcome_tracking.py index 9809edc2..a98091dd 100644 --- a/routstr/upstream/terminal_outcome_tracking.py +++ b/routstr/upstream/terminal_outcome_tracking.py @@ -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( diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 863e3834..acdc4aad 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -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() diff --git a/tests/unit/test_terminal_outcomes.py b/tests/unit/test_terminal_outcomes.py index b825a304..6a3dbf25 100644 --- a/tests/unit/test_terminal_outcomes.py +++ b/tests/unit/test_terminal_outcomes.py @@ -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: diff --git a/tests/unit/test_x_cashu_cost_sats.py b/tests/unit/test_x_cashu_cost_sats.py index 7b7873b8..bbdb385c 100644 --- a/tests/unit/test_x_cashu_cost_sats.py +++ b/tests/unit/test_x_cashu_cost_sats.py @@ -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) # ---------------------------------------------------------------------------