From 81c0ff57e92a0115a8045bdc8785fe61cc7d65a2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 24 Jul 2026 22:24:24 +0200 Subject: [PATCH] fix: expose cost breakdown in paid responses --- routstr/payment/cost_calculation.py | 30 +++- routstr/upstream/base.py | 140 +++++++++++++++++-- tests/unit/test_cost_calculation_caching.py | 4 + tests/unit/test_messages_litellm_dispatch.py | 3 +- tests/unit/test_x_cashu_cost_sats.py | 12 +- 5 files changed, 175 insertions(+), 14 deletions(-) diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 37ac15d3..d89f70ee 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -481,12 +481,38 @@ def _calculate_from_usd_cost( ) output_msats = cost_in_msats - input_msats + # Estimate cache read/creation msats proportionally within the input cost. + # These are informational subcomponents: input_msats remains inclusive of + # cache cost so input_msats + output_msats == total_msats, matching the + # token-priced path and the public CostData contract. + cache_read_msats = 0 + cache_creation_msats = 0 + if cache_read_tokens > 0 or cache_creation_tokens > 0: + cache_tokens = cache_read_tokens + cache_creation_tokens + regular_input_tokens = input_tokens + total_input_tokens = regular_input_tokens + cache_tokens + if total_input_tokens > 0: + # Approximate by token count because the USD path only exposes an + # aggregate input cost, not separately priced cache buckets. + cache_read_msats = ( + int(input_msats * cache_read_tokens / total_input_tokens) + if cache_read_tokens > 0 + else 0 + ) + cache_creation_msats = ( + int(input_msats * cache_creation_tokens / total_input_tokens) + if cache_creation_tokens > 0 + else 0 + ) + logger.info( "Using cost from usage data/details", extra={ "usd_cost": usd_cost, "cost_in_sats": cost_in_sats, "cost_in_msats": cost_in_msats, + "cache_read_msats": cache_read_msats, + "cache_creation_msats": cache_creation_msats, "model": response_data.get("model", "unknown"), }, ) @@ -501,8 +527,8 @@ def _calculate_from_usd_cost( output_tokens=output_tokens, cache_read_input_tokens=cache_read_tokens, cache_creation_input_tokens=cache_creation_tokens, - cache_read_msats=0, - cache_creation_msats=0, + cache_read_msats=cache_read_msats, + cache_creation_msats=cache_creation_msats, ) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index a8dba7e3..343230cf 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -69,6 +69,56 @@ if typing.TYPE_CHECKING: logger = get_logger(__name__) +def _inject_cost_response_headers( + headers: dict[str, str], cost_data: CostData | MaxCostData +) -> None: + """Inject per-request cost breakdown into response headers. + + The SDK's ``extractUsageFromResponseHeaders`` reads these to populate + ``inputMsats``, ``outputMsats``, ``totalMsats`` and ``satsCost`` in the + usage tracking entry — without them, x-cashu requests show 0.0 for all + sat cost fields. + """ + headers["X-Routstr-Cost-Msats"] = str(cost_data.total_msats) + headers["X-Routstr-Input-Cost-Msats"] = str(cost_data.input_msats) + headers["X-Routstr-Output-Cost-Msats"] = str(cost_data.output_msats) + if cost_data.total_usd: + headers["X-Routstr-Cost-Usd"] = str(cost_data.total_usd) + + +def _inject_cost_into_usage( + response_json: dict, cost_data: CostData | MaxCostData +) -> None: + """Inject cost breakdown into the response body's ``usage.cost`` object. + + The SDK's ``extractUsageFromResponseBody`` expects ``usage.cost`` to be + an object with ``total_msats``/``input_msats``/``output_msats`` (not a + plain USD number). When the upstream returns ``cost`` as a number, the + SDK cannot extract the msats breakdown from the body alone. + """ + usage = response_json.get("usage") + if not isinstance(usage, dict): + return + # Direct assignment (not setdefault) so routstr's authoritative cost + # data always overwrites any upstream-provided cost values. Using + # setdefault would silently keep stale upstream values and drop our + # calculated msats breakdown. + cost_obj: dict[str, int | float] = { + "base_msats": cost_data.base_msats, + "input_msats": cost_data.input_msats, + "output_msats": cost_data.output_msats, + "total_msats": cost_data.total_msats, + "cache_read_input_tokens": cost_data.cache_read_input_tokens, + "cache_creation_input_tokens": cost_data.cache_creation_input_tokens, + "cache_read_msats": cost_data.cache_read_msats, + "cache_creation_msats": cost_data.cache_creation_msats, + } + if cost_data.total_usd: + cost_obj["total_usd"] = cost_data.total_usd + usage["cost"] = cost_obj + usage["cost_sats"] = cost_data.total_msats // 1000 + + def _is_json_content_type(content_type: str | None) -> bool: """Return True when the upstream response should be parsed as JSON.""" if not content_type: @@ -284,9 +334,27 @@ class BaseUpstreamProvider: sats_cost = total_msats // 1000 + # Build the cost object that the SDK's extractUsageFromResponseBody + # and extractUsageFromSSEJson expect: an object with total_msats, + # input_msats, output_msats, cache_read_msats, cache_creation_msats, + # etc. Setting usage.cost to a plain float (total_usd) means the SDK + # cannot extract the msats breakdown — cache_read_msats and + # cache_creation_msats in particular are lost. + cost_obj = { + "base_msats": cost_dict.get("base_msats", 0), + "input_msats": cost_dict.get("input_msats", 0), + "output_msats": cost_dict.get("output_msats", 0), + "total_msats": total_msats, + "total_usd": total_usd, + "cache_read_input_tokens": cost_dict.get("cache_read_input_tokens", 0), + "cache_creation_input_tokens": cost_dict.get("cache_creation_input_tokens", 0), + "cache_read_msats": cost_dict.get("cache_read_msats", 0), + "cache_creation_msats": cost_dict.get("cache_creation_msats", 0), + } + # Inject into top-level usage block (OpenAI/Anthropic style) if "usage" in response_json: - response_json["usage"]["cost"] = total_usd + response_json["usage"]["cost"] = cost_obj response_json["usage"]["cost_sats"] = sats_cost response_json["usage"]["remaining_balance_msats"] = key.balance self._fold_cache_into_input_tokens(response_json["usage"]) @@ -2153,6 +2221,21 @@ class BaseUpstreamProvider: if k.lower() in allowed_headers } + # Inject cost breakdown headers so the SDK's + # extractUsageFromResponseHeaders can populate + # inputMsats/outputMsats/totalMsats for balance-mode requests. + if isinstance(cost_data, dict): + _cost_data_obj = CostData( + base_msats=cost_data.get("base_msats", 0), + input_msats=cost_data.get("input_msats", 0), + output_msats=cost_data.get("output_msats", 0), + total_msats=cost_data.get("total_msats", 0), + total_usd=cost_data.get("total_usd", 0.0), + ) + else: + _cost_data_obj = cost_data + _inject_cost_response_headers(response_headers, _cost_data_obj) + return Response( content=json.dumps(response_json).encode(), status_code=response.status_code, @@ -2242,9 +2325,24 @@ class BaseUpstreamProvider: ) self.inject_cost_metadata(response_json, cost_data, key) + # Inject cost breakdown headers for balance-mode requests. + if isinstance(cost_data, dict): + _cost_data_obj = CostData( + base_msats=cost_data.get("base_msats", 0), + input_msats=cost_data.get("input_msats", 0), + output_msats=cost_data.get("output_msats", 0), + total_msats=cost_data.get("total_msats", 0), + total_usd=cost_data.get("total_usd", 0.0), + ) + else: + _cost_data_obj = cost_data + response_headers: dict[str, str] = {} + _inject_cost_response_headers(response_headers, _cost_data_obj) + return Response( content=json.dumps(response_json).encode(), status_code=200, + headers=response_headers, media_type="application/json", ) @@ -2295,11 +2393,12 @@ class BaseUpstreamProvider: and "usage" in response_json and isinstance(response_json["usage"], dict) ): - response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + _inject_cost_into_usage(response_json, cost_data) self._fold_cache_into_input_tokens(response_json["usage"]) response_headers: dict[str, str] = {} if cost_data: + _inject_cost_response_headers(response_headers, cost_data) refund_amount = messages_dispatch.compute_refund( amount, unit, cost_data.total_msats ) @@ -3700,6 +3799,11 @@ class BaseUpstreamProvider: "model": model, }, ) + + # Inject cost breakdown headers so the SDK's + # extractUsageFromResponseHeaders can populate + # inputMsats/outputMsats/totalMsats for x-cashu requests. + _inject_cost_response_headers(response_headers, cost_data) except Exception as e: logger.error( "Error calculating cost for streaming response", @@ -3722,8 +3826,12 @@ class BaseUpstreamProvider: if "provider" not in data_json: self._apply_provider_field(data_json) changed = True - if cost_data and "usage" in data_json and data_json["usage"]: - data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + if ( + cost_data + and "usage" in data_json + and data_json["usage"] + ): + _inject_cost_into_usage(data_json, cost_data) changed = True if changed: lines[i] = "data: " + json.dumps(data_json) @@ -3777,7 +3885,10 @@ class BaseUpstreamProvider: ) if cost_data and "usage" in response_json: - response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + # Inject cost breakdown into both the response body (so the + # SDK's body extractor picks up the msats breakdown) and the + # response headers (so the SDK's header extractor works too). + _inject_cost_into_usage(response_json, cost_data) if not cost_data: logger.error( @@ -3808,6 +3919,8 @@ class BaseUpstreamProvider: if "content-encoding" in response_headers: del response_headers["content-encoding"] + _inject_cost_response_headers(response_headers, cost_data) + if unit == "msat": refund_amount = amount - cost_data.total_msats elif unit == "sat": @@ -4681,6 +4794,11 @@ class BaseUpstreamProvider: "model": model, }, ) + + # Inject cost breakdown headers so the SDK's + # extractUsageFromResponseHeaders can populate + # inputMsats/outputMsats/totalMsats for x-cashu requests. + _inject_cost_response_headers(response_headers, cost_data) except Exception as e: logger.error( "Error calculating cost for streaming Responses API response", @@ -4703,8 +4821,12 @@ class BaseUpstreamProvider: if "provider" not in data_json: self._apply_provider_field(data_json) changed = True - if cost_data and "usage" in data_json and data_json["usage"]: - data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + if ( + cost_data + and "usage" in data_json + and data_json["usage"] + ): + _inject_cost_into_usage(data_json, cost_data) changed = True if changed: lines[i] = "data: " + json.dumps(data_json) @@ -4747,7 +4869,7 @@ class BaseUpstreamProvider: ) if cost_data and "usage" in response_json: - response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + _inject_cost_into_usage(response_json, cost_data) if not cost_data: logger.error( @@ -4778,6 +4900,8 @@ class BaseUpstreamProvider: if "content-encoding" in response_headers: del response_headers["content-encoding"] + _inject_cost_response_headers(response_headers, cost_data) + if unit == "msat": refund_amount = amount - cost_data.total_msats elif unit == "sat": diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index ba31366a..0722e5bf 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -529,6 +529,8 @@ async def test_openrouter_upstream_inference_cost_components_are_used() -> None: assert isinstance(result, CostData) assert result.input_msats == 994 assert result.output_msats == 3477 + assert result.cache_read_msats == 758 + assert result.cache_creation_msats == 0 assert result.input_msats + result.output_msats == result.total_msats == 4471 @@ -574,6 +576,8 @@ async def test_ppq_byok_bills_upstream_inference_cost_plus_fee() -> None: # Token normalisation (OpenAI dialect: cached included in prompt_tokens) assert result.input_tokens == 5070 # 164371 - 159301 assert result.cache_read_input_tokens == 159301 + assert result.cache_read_msats == 897966 + assert result.cache_creation_msats == 0 assert result.output_tokens == 99 diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index c41568ad..79d16030 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -451,7 +451,8 @@ async def test_non_streaming_dispatches_via_litellm_and_returns_anthropic_respon assert payload["model"] == "openai/gpt-4o-mini" # mapped back to requested assert payload["usage"]["input_tokens"] == 5 assert payload["usage"]["output_tokens"] == 3 - assert payload["usage"]["cost"] == 0.0001 + assert payload["usage"]["cost"]["total_msats"] == 1234 + assert payload["usage"]["cost"]["total_usd"] == 0.0001 assert payload["usage"]["cost_sats"] == 1 diff --git a/tests/unit/test_x_cashu_cost_sats.py b/tests/unit/test_x_cashu_cost_sats.py index 0dc509cf..901cc2f6 100644 --- a/tests/unit/test_x_cashu_cost_sats.py +++ b/tests/unit/test_x_cashu_cost_sats.py @@ -67,8 +67,13 @@ async def test_non_streaming_includes_cost_sats() -> None: ) body = json.loads(response.body) - assert "cost_sats" in body["usage"] assert body["usage"]["cost_sats"] == 5 # 5000 msats // 1000 + assert body["usage"]["cost"]["total_msats"] == 5000 + assert body["usage"]["cost"]["input_msats"] == 3000 + assert body["usage"]["cost"]["output_msats"] == 2000 + assert response.headers["x-routstr-cost-msats"] == "5000" + assert response.headers["x-routstr-input-cost-msats"] == "3000" + assert response.headers["x-routstr-output-cost-msats"] == "2000" @pytest.mark.asyncio @@ -96,7 +101,7 @@ async def test_non_streaming_cost_sats_value_rounds_down() -> None: @pytest.mark.asyncio -async def test_non_streaming_preserves_existing_usage_fields() -> None: +async def test_non_streaming_preserves_tokens_and_replaces_upstream_cost() -> None: provider = _make_provider() cost_data = _make_cost_data(total_msats=3000) @@ -127,7 +132,8 @@ async def test_non_streaming_preserves_existing_usage_fields() -> None: assert usage["prompt_tokens"] == 100 assert usage["completion_tokens"] == 50 assert usage["total_tokens"] == 150 - assert usage["cost"] == 0.00015 + assert usage["cost"]["total_msats"] == 3000 + assert usage["cost"]["total_usd"] == 0.00025 assert usage["cost_sats"] == 3