From 81c0ff57e92a0115a8045bdc8785fe61cc7d65a2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 24 Jul 2026 22:24:24 +0200 Subject: [PATCH 1/2] 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 From 4c6bc49e072327bf9001d64f4f216476e18ece6a Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 25 Jul 2026 00:36:32 +0200 Subject: [PATCH 2/2] unify payment path --- routstr/payment/cost_calculation.py | 81 ++++++--- routstr/upstream/base.py | 181 +++++++++---------- tests/unit/test_cost_calculation_caching.py | 132 +++++++++++++- tests/unit/test_cost_response_metadata.py | 113 ++++++++++++ tests/unit/test_messages_litellm_dispatch.py | 6 + 5 files changed, 395 insertions(+), 118 deletions(-) create mode 100644 tests/unit/test_cost_response_metadata.py diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index d89f70ee..0fb305cb 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -183,6 +183,28 @@ async def calculate_cost( cost_details.get("output_cost") or cost_details.get("upstream_inference_completions_cost") ) + cache_pricing_rates: tuple[float, float, float, float] | None = None + if cache_read_tokens > 0 or cache_creation_tokens > 0: + try: + cache_pricing_rates = _get_pricing_rates( + response_data, model_obj, provider_fee + ) + except ValueError: + logger.warning( + "Cache pricing unavailable for USD cost breakdown; " + "leaving cache cost components unknown", + extra={"model": response_data.get("model", "unknown")}, + ) + if cache_pricing_rates is None and settings.fixed_pricing: + fixed_input_rate = ( + float(settings.fixed_per_1k_input_tokens) * 1000.0 + ) + cache_pricing_rates = ( + fixed_input_rate, + float(settings.fixed_per_1k_output_tokens) * 1000.0, + fixed_input_rate, + fixed_input_rate, + ) return _calculate_from_usd_cost( usd_cost, input_usd, @@ -193,6 +215,7 @@ async def calculate_cost( output_tokens, response_data, provider_fee, + cache_pricing_rates, ) except Exception as e: logger.warning( @@ -451,6 +474,7 @@ def _calculate_from_usd_cost( output_tokens: int, response_data: dict, provider_fee: float | None, + pricing_rates: tuple[float, float, float, float] | None = None, ) -> CostData: """Calculate cost from USD figures, deriving input/output split from tokens.""" if provider_fee is None: @@ -460,15 +484,20 @@ def _calculate_from_usd_cost( output_usd = output_usd * provider_fee sats_per_usd = 1.0 / sats_usd_price() cost_in_sats = usd_cost * sats_per_usd - cost_in_msats = math.ceil(cost_in_sats * 1000) + raw_cost_msats = cost_in_sats * 1000 + cost_in_msats = math.ceil(raw_cost_msats) + raw_input_msats = 0.0 if input_usd > 0 or output_usd > 0: # The total is the authoritative billed amount. Allocating that integer # total proportionally avoids losing sub-millisatoshi remainders when # input and output components are each truncated independently. component_usd = input_usd + output_usd - input_msats = math.floor(cost_in_msats * input_usd / component_usd) - output_msats = cost_in_msats - input_msats + # Match the token-priced path: truncate the visible output component + # and assign the authoritative total's rounding remainder to input. + output_msats = math.floor(cost_in_msats * output_usd / component_usd) + input_msats = cost_in_msats - output_msats + raw_input_msats = raw_cost_msats * input_usd / component_usd else: effective_input_tokens = ( input_tokens + cache_read_tokens + cache_creation_tokens @@ -480,29 +509,37 @@ def _calculate_from_usd_cost( else 0 ) output_msats = cost_in_msats - input_msats + raw_input_msats = ( + raw_cost_msats * effective_input_tokens / total_tokens + if total_tokens > 0 + else 0.0 + ) - # 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. + # Preserve the same cache-rate ratios as the token-priced path while the + # upstream USD total remains authoritative. Cache values are informational + # subcomponents of the inclusive input cost. 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 + if pricing_rates is not None: + input_rate, _, cache_read_rate, cache_creation_rate = pricing_rates + regular_weight = input_tokens * input_rate + cache_read_weight = cache_read_tokens * cache_read_rate + cache_creation_weight = cache_creation_tokens * cache_creation_rate + total_input_weight = ( + regular_weight + cache_read_weight + cache_creation_weight + ) + if total_input_weight > 0: + cache_read_msats = int( + round( + raw_input_msats * cache_read_weight / total_input_weight, + 3, + ) ) - cache_creation_msats = ( - int(input_msats * cache_creation_tokens / total_input_tokens) - if cache_creation_tokens > 0 - else 0 + cache_creation_msats = int( + round( + raw_input_msats * cache_creation_weight / total_input_weight, + 3, + ) ) logger.info( diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 343230cf..3353df8a 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -69,8 +69,21 @@ if typing.TYPE_CHECKING: logger = get_logger(__name__) +CostMetadata = CostData | MaxCostData | dict[str, Any] + + +def _cost_field( + cost_data: CostMetadata, field: str, default: int | float = 0 +) -> int | float: + if isinstance(cost_data, dict): + value = cost_data.get(field, default) + else: + value = getattr(cost_data, field, default) + return value if isinstance(value, (int, float)) else default + + def _inject_cost_response_headers( - headers: dict[str, str], cost_data: CostData | MaxCostData + headers: dict[str, str], cost_data: CostMetadata ) -> None: """Inject per-request cost breakdown into response headers. @@ -79,16 +92,21 @@ def _inject_cost_response_headers( 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) + headers["X-Routstr-Cost-Msats"] = str( + int(_cost_field(cost_data, "total_msats")) + ) + headers["X-Routstr-Input-Cost-Msats"] = str( + int(_cost_field(cost_data, "input_msats")) + ) + headers["X-Routstr-Output-Cost-Msats"] = str( + int(_cost_field(cost_data, "output_msats")) + ) + total_usd = float(_cost_field(cost_data, "total_usd", 0.0)) + if total_usd: + headers["X-Routstr-Cost-Usd"] = str(total_usd) -def _inject_cost_into_usage( - response_json: dict, cost_data: CostData | MaxCostData -) -> None: +def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> None: """Inject cost breakdown into the response body's ``usage.cost`` object. The SDK's ``extractUsageFromResponseBody`` expects ``usage.cost`` to be @@ -104,19 +122,26 @@ def _inject_cost_into_usage( # 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, + "base_msats": int(_cost_field(cost_data, "base_msats")), + "input_msats": int(_cost_field(cost_data, "input_msats")), + "output_msats": int(_cost_field(cost_data, "output_msats")), + "total_msats": int(_cost_field(cost_data, "total_msats")), + "cache_read_input_tokens": int( + _cost_field(cost_data, "cache_read_input_tokens") + ), + "cache_creation_input_tokens": int( + _cost_field(cost_data, "cache_creation_input_tokens") + ), + "cache_read_msats": int(_cost_field(cost_data, "cache_read_msats")), + "cache_creation_msats": int( + _cost_field(cost_data, "cache_creation_msats") + ), } - if cost_data.total_usd: - cost_obj["total_usd"] = cost_data.total_usd + total_usd = float(_cost_field(cost_data, "total_usd", 0.0)) + if total_usd: + cost_obj["total_usd"] = total_usd usage["cost"] = cost_obj - usage["cost_sats"] = cost_data.total_msats // 1000 + usage["cost_sats"] = int(_cost_field(cost_data, "total_msats")) // 1000 def _is_json_content_type(content_type: str | None) -> bool: @@ -325,48 +350,24 @@ class BaseUpstreamProvider: self._apply_provider_field(response_json) if isinstance(cost_data, dict): total_msats = cost_data.get("total_msats", 0) - total_usd = cost_data.get("total_usd", 0.0) cost_dict = cost_data else: total_msats = cost_data.total_msats - total_usd = cost_data.total_usd cost_dict = cost_data.dict() 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"] = cost_obj - response_json["usage"]["cost_sats"] = sats_cost + # Inject the shared SDK cost contract into every usage shape. + if isinstance(response_json.get("usage"), dict): + _inject_cost_into_usage(response_json, cost_data) response_json["usage"]["remaining_balance_msats"] = key.balance self._fold_cache_into_input_tokens(response_json["usage"]) - # Inject into Anthropic nested usage block if present - if ( - "message" in response_json - and isinstance(response_json["message"], dict) - and "usage" in response_json["message"] - ): - response_json["message"]["usage"]["sats_cost"] = sats_cost - self._fold_cache_into_input_tokens(response_json["message"]["usage"]) + message = response_json.get("message") + if isinstance(message, dict) and isinstance(message.get("usage"), dict): + _inject_cost_into_usage(message, cost_data) + message["usage"]["remaining_balance_msats"] = key.balance + self._fold_cache_into_input_tokens(message["usage"]) # Unified Routstr metadata response_json["metadata"] = response_json.get("metadata", {}) @@ -1297,12 +1298,9 @@ class BaseUpstreamProvider: await session.refresh(key) remaining_balance_msats = key.balance - # Merge cost into usage for OpenCode + # Merge the shared cost contract into usage for SDKs and OpenCode. if "usage" in response_json: - response_json["usage"]["cost"] = cost_data.get("total_usd", 0.0) - response_json["usage"]["cost_sats"] = ( - cost_data.get("total_msats", 0) // 1000 - ) + _inject_cost_into_usage(response_json, cost_data) response_json["usage"]["remaining_balance_msats"] = ( remaining_balance_msats ) @@ -1349,6 +1347,7 @@ class BaseUpstreamProvider: for k, v in response.headers.items() if k.lower() in allowed_headers } + _inject_cost_response_headers(response_headers, cost_data) if requested_model: response_json["model"] = requested_model @@ -1734,12 +1733,9 @@ class BaseUpstreamProvider: await session.refresh(key) remaining_balance_msats = key.balance - # Merge cost into usage for OpenCode + # Merge the shared cost contract into usage for SDKs and OpenCode. if "usage" in response_json: - response_json["usage"]["cost"] = cost_data.get("total_usd", 0.0) - response_json["usage"]["cost_sats"] = ( - cost_data.get("total_msats", 0) // 1000 - ) + _inject_cost_into_usage(response_json, cost_data) response_json["usage"]["remaining_balance_msats"] = ( remaining_balance_msats ) @@ -1786,6 +1782,7 @@ class BaseUpstreamProvider: for k, v in response.headers.items() if k.lower() in allowed_headers } + _inject_cost_response_headers(response_headers, cost_data) if requested_model: response_json["model"] = requested_model @@ -2221,20 +2218,8 @@ 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) + # Inject the same cost headers used by every paid response path. + _inject_cost_response_headers(response_headers, cost_data) return Response( content=json.dumps(response_json).encode(), @@ -2325,19 +2310,9 @@ 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 + # Inject the same cost headers used by every paid response path. response_headers: dict[str, str] = {} - _inject_cost_response_headers(response_headers, _cost_data_obj) + _inject_cost_response_headers(response_headers, cost_data) return Response( content=json.dumps(response_json).encode(), @@ -2644,7 +2619,7 @@ class BaseUpstreamProvider: the cost of a wire-format change for clients that read ``X-Cashu`` from headers today. """ - buffered: list[bytes] = [] + buffered: list[messages_dispatch.AnnotatedEvent] = [] last_model_seen: str | None = None input_tokens = 0 output_tokens = 0 @@ -2672,7 +2647,7 @@ class BaseUpstreamProvider: total_cost = max(total_cost, annotated.total_cost) input_cost = max(input_cost, annotated.input_cost) output_cost = max(output_cost, annotated.output_cost) - buffered.append(annotated.sse_bytes) + buffered.append(annotated) response_headers: dict[str, str] = { "Cache-Control": "no-cache", @@ -2699,6 +2674,7 @@ class BaseUpstreamProvider: }, ) + cost_data: CostData | MaxCostData | None = None if ( input_tokens > 0 or output_tokens > 0 @@ -2754,9 +2730,30 @@ class BaseUpstreamProvider: }, ) + if cost_data: + _inject_cost_response_headers(response_headers, cost_data) + for index, annotated in enumerate(buffered): + event = annotated.event + changed = False + message = event.get("message") + if isinstance(message, dict) and isinstance(message.get("usage"), dict): + _inject_cost_into_usage(message, cost_data) + changed = True + if isinstance(event.get("usage"), dict): + _inject_cost_into_usage(event, cost_data) + changed = True + if changed: + event_type = str(event.get("type") or "") + prefix = f"event: {event_type}\n" if event_type else "" + buffered[index] = annotated._replace( + sse_bytes=( + f"{prefix}data: {json.dumps(event)}\n\n".encode() + ) + ) + async def replay() -> AsyncGenerator[bytes, None]: - for chunk in buffered: - yield chunk + for annotated in buffered: + yield annotated.sse_bytes return StreamingResponse( replay(), diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index 0722e5bf..65ad5091 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -15,6 +15,7 @@ os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to") from routstr.core.settings import settings from routstr.payment.cost_calculation import CostData, MaxCostData, calculate_cost +from routstr.payment.models import Architecture, Model, Pricing @pytest.fixture(autouse=True) @@ -527,13 +528,136 @@ async def test_openrouter_upstream_inference_cost_components_are_used() -> None: result = await calculate_cost(response, max_cost=100000) assert isinstance(result, CostData) - assert result.input_msats == 994 - assert result.output_msats == 3477 + assert result.input_msats == 995 + assert result.output_msats == 3476 assert result.cache_read_msats == 758 assert result.cache_creation_msats == 0 assert result.input_msats + result.output_msats == result.total_msats == 4471 +@pytest.mark.asyncio +async def test_usd_cache_breakdown_matches_token_priced_path( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Authoritative USD totals must retain model-specific cache-rate ratios.""" + monkeypatch.setattr(settings, "fixed_pricing", False) + model = Model( + id="cache-priced-model", + name="cache-priced-model", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="test", + instruct_type=None, + ), + pricing=Pricing(prompt=0.01, completion=0.02), + sats_pricing=Pricing( + prompt=0.01, + completion=0.02, + input_cache_read=0.001, + input_cache_write=0.01, + ), + per_request_limits=None, + top_provider=None, + ) + usage = { + "prompt_tokens": 1000, + "completion_tokens": 100, + "prompt_tokens_details": {"cached_tokens": 900}, + } + + token_result = await calculate_cost( + {"model": model.id, "usage": usage}, + max_cost=100_000, + model_obj=model, + ) + usd_result = await calculate_cost( + { + "model": model.id, + "usage": { + **usage, + "cost": 0.000195, + "cost_details": { + "input_cost": 0.000095, + "output_cost": 0.0001, + }, + }, + }, + max_cost=100_000, + model_obj=model, + provider_fee=1.0, + ) + + assert isinstance(token_result, CostData) + assert isinstance(usd_result, CostData) + assert usd_result.total_msats == token_result.total_msats == 3900 + assert usd_result.input_msats + usd_result.output_msats == usd_result.total_msats + assert usd_result.cache_read_msats == token_result.cache_read_msats == 900 + + +@pytest.mark.asyncio +async def test_usd_cache_breakdown_does_not_absorb_total_rounding_remainder( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Sub-msat cache components truncate like the token-priced path.""" + monkeypatch.setattr(settings, "fixed_pricing", False) + model = Model( + id="sub-msat-cache-model", + name="sub-msat-cache-model", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="test", + instruct_type=None, + ), + pricing=Pricing(prompt=0.001, completion=0.001), + sats_pricing=Pricing( + prompt=0.001, + completion=0.001, + input_cache_write=0.0006, + ), + per_request_limits=None, + top_provider=None, + ) + usage = { + "input_tokens": 0, + "output_tokens": 0, + "cache_creation_input_tokens": 1, + } + + token_result = await calculate_cost( + {"model": model.id, "usage": usage}, + max_cost=100_000, + model_obj=model, + ) + usd_result = await calculate_cost( + { + "model": model.id, + "usage": { + **usage, + "cost": 0.00000003, + "cost_details": {"input_cost": 0.00000003}, + }, + }, + max_cost=100_000, + model_obj=model, + provider_fee=1.0, + ) + + assert isinstance(token_result, CostData) + assert isinstance(usd_result, CostData) + assert usd_result.total_msats == token_result.total_msats == 1 + assert usd_result.cache_creation_msats == token_result.cache_creation_msats == 0 + + # ============================================================================ # PPQ.AI BYOK: upstream_inference_cost + BYOK fee billing # @@ -570,8 +694,8 @@ async def test_ppq_byok_bills_upstream_inference_cost_plus_fee() -> None: # msats), not the fee alone (~0.0023 USD → ~45k msats). ~20× correction. assert result.total_msats == 940274 assert result.input_msats + result.output_msats == result.total_msats - assert result.input_msats == 926546 - assert result.output_msats == 13728 + assert result.input_msats == 926547 + assert result.output_msats == 13727 assert result.total_usd == pytest.approx(0.047013667305) # Token normalisation (OpenAI dialect: cached included in prompt_tokens) assert result.input_tokens == 5070 # 164371 - 159301 diff --git a/tests/unit/test_cost_response_metadata.py b/tests/unit/test_cost_response_metadata.py new file mode 100644 index 00000000..beaf4005 --- /dev/null +++ b/tests/unit/test_cost_response_metadata.py @@ -0,0 +1,113 @@ +"""Response-contract tests for Routstr cost metadata across paid paths.""" + +import json +import os +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.core.db import ApiKey # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 + +COST_DATA = { + "base_msats": 0, + "input_msats": 1_200, + "output_msats": 300, + "total_msats": 1_500, + "total_usd": 0.0001, + "input_tokens": 10, + "output_tokens": 3, + "cache_read_input_tokens": 8, + "cache_creation_input_tokens": 2, + "cache_read_msats": 80, + "cache_creation_msats": 40, +} + + +def _provider() -> BaseUpstreamProvider: + return BaseUpstreamProvider(base_url="http://test", api_key="upstream-key") + + +def _key() -> ApiKey: + return ApiKey(hashed_key="abcdef0123" * 4, balance=1_000_000) + + +def _session() -> Any: + session = MagicMock() + session.refresh = AsyncMock() + return session + + +def _upstream_response(payload: dict) -> httpx.Response: + return httpx.Response( + 200, + json=payload, + request=httpx.Request("POST", "http://test"), + ) + + +def _assert_cost_contract(response: Any) -> None: + body = json.loads(response.body) + assert body["usage"]["cost"] == { + "base_msats": 0, + "input_msats": 1_200, + "output_msats": 300, + "total_msats": 1_500, + "total_usd": 0.0001, + "cache_read_input_tokens": 8, + "cache_creation_input_tokens": 2, + "cache_read_msats": 80, + "cache_creation_msats": 40, + } + assert response.headers["X-Routstr-Cost-Msats"] == "1500" + assert response.headers["X-Routstr-Input-Cost-Msats"] == "1200" + assert response.headers["X-Routstr-Output-Cost-Msats"] == "300" + + +@pytest.mark.asyncio +async def test_balance_chat_completion_uses_shared_cost_contract() -> None: + provider = _provider() + with patch( + "routstr.upstream.base.adjust_payment_for_tokens", + new=AsyncMock(return_value=dict(COST_DATA)), + ): + response = await provider.handle_non_streaming_chat_completion( + _upstream_response( + { + "model": "test-model", + "usage": {"prompt_tokens": 10, "completion_tokens": 3}, + } + ), + _key(), + _session(), + deducted_max_cost=10_000, + ) + + _assert_cost_contract(response) + + +@pytest.mark.asyncio +async def test_balance_responses_completion_uses_shared_cost_contract() -> None: + provider = _provider() + with patch( + "routstr.upstream.base.adjust_payment_for_tokens", + new=AsyncMock(return_value=dict(COST_DATA)), + ): + response = await provider.handle_non_streaming_responses_completion( + _upstream_response( + { + "model": "test-model", + "usage": {"input_tokens": 10, "output_tokens": 3}, + } + ), + _key(), + _session(), + deducted_max_cost=10_000, + ) + + _assert_cost_contract(response) diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index 79d16030..cb3fe9d7 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -855,6 +855,9 @@ async def test_x_cashu_streaming_replays_events_and_sets_refund_header() -> None assert isinstance(result, StreamingResponse) assert result.headers.get("X-Cashu") == "cashuSTREAM" + assert result.headers.get("X-Routstr-Cost-Msats") == "1500000" + assert result.headers.get("X-Routstr-Input-Cost-Msats") == "1000000" + assert result.headers.get("X-Routstr-Output-Cost-Msats") == "500000" # 1_500_000 msats → 1500 sats. Refund = 5000 - 1500 = 3500. mock_refund.assert_awaited_once() refund_call = mock_refund.await_args @@ -873,6 +876,9 @@ async def test_x_cashu_streaming_replays_events_and_sets_refund_header() -> None assert "event: message_start" in joined assert "event: message_delta" in joined assert "event: message_stop" in joined + assert '"total_msats": 1500000' in joined + assert '"input_msats": 1000000' in joined + assert '"output_msats": 500000' in joined # ---------------------------------------------------------------------------