refactor: share x-cashu usage estimation and log full refunds

This commit is contained in:
9qeklajc
2026-09-08 22:05:07 +02:00
parent f58983b9eb
commit 12befe7a3d
3 changed files with 145 additions and 52 deletions
+105 -52
View File
@@ -154,13 +154,33 @@ def _inject_cost_response_headers(
headers["X-Routstr-Cost-Usd"] = str(total_usd)
def _estimated_usage(
estimator: MissingUsageEstimator, model: str | None
) -> dict[str, Any] | None:
"""Local usage estimate, or None when the upstream generated no text."""
if not estimator.output_text:
return None
return estimator.response_data(model)["usage"]
def _apply_estimated_usage(
response_json: dict[str, Any],
request_body: bytes | None,
model_obj: Model | None,
amount: int,
unit: str,
api: str,
) -> None:
"""Bill a buffered response from a local estimate when usage is missing."""
if response_json.get("usage"):
return
estimator = MissingUsageEstimator(request_body, model_obj)
estimator.observe(response_json)
estimated = estimator.estimated_usage(response_json.get("model"))
if not estimated:
return
logger.warning(
"No usage in non-streaming response, billing from local token estimate",
extra={
"api": api,
"model": response_json.get("model", "unknown"),
"amount": amount,
"unit": unit,
"estimated_usage": estimated,
},
)
response_json["usage"] = estimated
def _parse_sse_events(content: str) -> list[tuple[list[str], str]]:
@@ -472,6 +492,34 @@ class BaseUpstreamProvider:
return
response_json["provider"] = f"{provider_type}:{existing_str}"
def _log_full_refund(
self,
*,
route: str,
model: str | None,
content_str: str,
amount: int,
unit: str,
) -> None:
"""Record a settlement that serves content but charges nothing.
The client keeps both the response and the whole prepayment, so the
model, the serving upstream and a redacted body preview are logged to
keep the unbilled request auditable.
"""
logger.warning(
"Zero-cost settlement, refunding the full prepayment",
extra={
"route": route,
"model": model or "unknown",
"provider_type": self.provider_type,
"upstream_base_url": self.base_url,
"refund_amount": amount,
"unit": unit,
"response_body_preview": redact_org_ids(content_str.strip()[:500]),
},
)
def inject_cost_metadata(
self,
response_json: dict,
@@ -3893,7 +3941,7 @@ class BaseUpstreamProvider:
usage_estimator.observe(data_json)
if not usage_data:
usage_data = _estimated_usage(usage_estimator, model)
usage_data = usage_estimator.estimated_usage(model)
if usage_data:
logger.warning(
"No usage in streaming response, billing from local token estimate",
@@ -3920,6 +3968,14 @@ class BaseUpstreamProvider:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model, model_obj
)
if cost_data is not None and cost_data.total_msats == 0:
self._log_full_refund(
route="chat.streaming",
model=model,
content_str=content_str,
amount=amount,
unit=unit,
)
if cost_data:
if unit == "msat":
refund_amount = amount - cost_data.total_msats
@@ -4044,26 +4100,20 @@ class BaseUpstreamProvider:
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
if not response_json.get("usage"):
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
estimated = _estimated_usage(
usage_estimator, response_json.get("model")
)
if estimated:
logger.warning(
"No usage in non-streaming response, billing from local token estimate",
extra={
"model": response_json.get("model", "unknown"),
"amount": amount,
"unit": unit,
"estimated_usage": estimated,
},
)
response_json["usage"] = estimated
_apply_estimated_usage(
response_json, request_body, model_obj, amount, unit, "chat"
)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data is not None and cost_data.total_msats == 0:
self._log_full_refund(
route="chat",
model=response_json.get("model"),
content_str=content_str,
amount=amount,
unit=unit,
)
if cost_data and "usage" in response_json:
# Inject cost breakdown into both the response body (so the
@@ -4936,16 +4986,17 @@ class BaseUpstreamProvider:
model = payload["model"]
if not usage_data:
usage_data = _estimated_usage(usage_estimator, model)
logger.warning(
"No usage in streaming Responses API response, billing from local token estimate",
extra={
"model": model,
"amount": amount,
"unit": unit,
"estimated_usage": usage_data,
},
)
usage_data = usage_estimator.estimated_usage(model)
if usage_data:
logger.warning(
"No usage in streaming Responses API response, billing from local token estimate",
extra={
"model": model,
"amount": amount,
"unit": unit,
"estimated_usage": usage_data,
},
)
else:
logger.debug(
"Found usage data in streaming Responses API response",
@@ -4963,6 +5014,14 @@ class BaseUpstreamProvider:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model, model_obj
)
if cost_data is not None and cost_data.total_msats == 0:
self._log_full_refund(
route="responses.streaming",
model=model,
content_str=content_str,
amount=amount,
unit=unit,
)
if cost_data:
if unit == "msat":
refund_amount = amount - cost_data.total_msats
@@ -5079,26 +5138,20 @@ class BaseUpstreamProvider:
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
if not response_json.get("usage"):
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
estimated = _estimated_usage(
usage_estimator, response_json.get("model")
)
if estimated:
logger.warning(
"No usage in non-streaming Responses API response, billing from local token estimate",
extra={
"model": response_json.get("model", "unknown"),
"amount": amount,
"unit": unit,
"estimated_usage": estimated,
},
)
response_json["usage"] = estimated
_apply_estimated_usage(
response_json, request_body, model_obj, amount, unit, "responses"
)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data is not None and cost_data.total_msats == 0:
self._log_full_refund(
route="responses",
model=response_json.get("model"),
content_str=content_str,
amount=amount,
unit=unit,
)
if cost_data and "usage" in response_json:
_inject_cost_into_usage(response_json, cost_data)
+6
View File
@@ -222,6 +222,12 @@ class MissingUsageEstimator:
return
self._output_parts.extend(_generated_text(response_data))
def estimated_usage(self, model: str | None = None) -> dict[str, Any] | None:
"""Local usage estimate, or None when the upstream generated no text."""
if not self.output_text:
return None
return self.response_data(model)["usage"]
def billing_data(
self,
response_data: dict[str, Any] | None,
+34
View File
@@ -5,6 +5,7 @@ When nothing can be estimated the prepayment is refunded in full.
"""
import json
import logging
import os
from typing import Any
from unittest.mock import AsyncMock, patch
@@ -267,3 +268,36 @@ async def test_pricing_error_refunds_full_prepayment(
assert result.status_code == 200
assert result.headers["X-Cashu"] == "cashuBrefund"
assert result.headers["X-Routstr-Cost-Msats"] == "0"
@pytest.mark.asyncio
@pytest.mark.parametrize("responses_api", [False, True])
async def test_full_refund_is_logged_with_model_provider_and_body(
responses_api: bool, caplog: pytest.LogCaptureFixture
) -> None:
payload = {
"model": "unpriced-model",
"usage": {"input_tokens": 10, "output_tokens": 5},
}
base_logger = logging.getLogger("routstr.upstream.base")
base_logger.addHandler(caplog.handler)
try:
with patch(
"routstr.payment.cost_calculation._get_pricing_rates",
side_effect=ValueError("No pricing for model"),
):
await _settle(_json(payload), responses_api=responses_api)
finally:
base_logger.removeHandler(caplog.handler)
record = next(
r
for r in caplog.records
if r.getMessage() == "Zero-cost settlement, refunding the full prepayment"
)
assert record.model == "unpriced-model"
assert record.provider_type == "base"
assert record.upstream_base_url == "http://test"
assert record.refund_amount == 10_000
assert record.unit == "msat"
assert "unpriced-model" in record.response_body_preview