diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index d41b2664..1f018f4e 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -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) diff --git a/routstr/upstream/count_tokens.py b/routstr/upstream/count_tokens.py index 51ffbfa4..ea067dd1 100644 --- a/routstr/upstream/count_tokens.py +++ b/routstr/upstream/count_tokens.py @@ -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, diff --git a/tests/unit/test_x_cashu_missing_usage.py b/tests/unit/test_x_cashu_missing_usage.py index ceaabade..fbeaa6e4 100644 --- a/tests/unit/test_x_cashu_missing_usage.py +++ b/tests/unit/test_x_cashu_missing_usage.py @@ -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