mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
refactor: share x-cashu usage estimation and log full refunds
This commit is contained in:
+105
-52
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user