mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: take reported usage per count instead of merging usage dialects
This commit is contained in:
+1
-1
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user