From 0aebfc6dbe9687e67c163c8b5121774cb12f62c5 Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Fri, 17 Jul 2026 16:01:16 +0200 Subject: [PATCH] fix: bill the serving provider's fee on the USD-cost path MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The USD-cost path (and the litellm pricing fallback) resolved the provider fee via get_provider_for_model(model_id)[0] — the best-ranked provider for the alias, not the one that served. Settlement callers in the upstream handlers now pass their own provider_fee through adjust_payment_for_tokens / get_x_cashu_cost into calculate_cost; the string-derived fallback remains for callers without a serving provider. Configured model pricing is unaffected (the fee is already baked into cached pricing). Co-Authored-By: Claude Fable 5 --- routstr/auth.py | 5 ++- routstr/payment/cost_calculation.py | 17 +++++++-- routstr/upstream/base.py | 13 +++++++ tests/unit/test_messages_litellm_dispatch.py | 14 ++++++- tests/unit/test_settlement_identity.py | 39 ++++++++++++++++++++ 5 files changed, 82 insertions(+), 6 deletions(-) diff --git a/routstr/auth.py b/routstr/auth.py index 516ef29c..3dcd8660 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -775,6 +775,7 @@ async def adjust_payment_for_tokens( session: AsyncSession, deducted_max_cost: int, model_obj: "Model | None" = None, + provider_fee: float | None = None, ) -> dict: """ Adjusts the payment based on token usage in the response. @@ -869,7 +870,9 @@ async def adjust_payment_for_tokens( extra={"error": str(e), "fee_msats": fee_msats}, ) - match await calculate_cost(response_data, deducted_max_cost, model_obj): + match await calculate_cost( + response_data, deducted_max_cost, model_obj, provider_fee + ): case MaxCostData() as cost: logger.debug( "Using max cost data (no token adjustment)", diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index aed83bd0..b9d03d22 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -71,6 +71,7 @@ async def calculate_cost( response_data: dict, max_cost: int, model_obj: "Model | None" = None, + provider_fee: float | None = None, ) -> CostData | MaxCostData | CostDataError: """Calculate the cost of an API request based on token usage. @@ -81,6 +82,11 @@ async def calculate_cost( its pricing is billed directly; without it, pricing is re-derived from the response's model string via the alias map, which resolves to the best-ranked candidate — not necessarily the serving one. + provider_fee: The serving provider's fee multiplier, applied on the + USD-cost path and the litellm pricing fallback (configured model + pricing already carries the fee baked in). Without it, the fee is + re-derived from the response's model string, which yields the + best-ranked provider's fee. Returns: Cost data or error information @@ -186,6 +192,7 @@ async def calculate_cost( cache_creation_tokens, output_tokens, response_data, + provider_fee, ) except Exception as e: logger.warning( @@ -199,7 +206,7 @@ async def calculate_cost( # Fall back to token-based pricing try: - pricing_rates = _get_pricing_rates(response_data, model_obj) + pricing_rates = _get_pricing_rates(response_data, model_obj, provider_fee) except ValueError as e: return CostDataError(message=str(e), code="pricing_error") @@ -317,6 +324,7 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: def _get_pricing_rates( response_data: dict, model_obj: "Model | None" = None, + provider_fee: float | None = None, ) -> tuple[float, float, float, float] | None: """Get configured rates, falling back to LiteLLM's model cost map. @@ -380,7 +388,8 @@ def _get_pricing_rates( if input_usd <= 0 or output_usd <= 0: raise ValueError(f"Incomplete LiteLLM pricing for model: {pricing_model}") - provider_fee = _resolve_provider_fee(response_model) + if provider_fee is None: + provider_fee = _resolve_provider_fee(response_model) usd_per_sat = sats_usd_price() mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat @@ -441,9 +450,11 @@ def _calculate_from_usd_cost( cache_creation_tokens: int, output_tokens: int, response_data: dict, + provider_fee: float | None = None, ) -> CostData: """Calculate cost from USD figures, deriving input/output split from tokens.""" - provider_fee = _resolve_provider_fee(response_data.get("model", "")) + if provider_fee is None: + provider_fee = _resolve_provider_fee(response_data.get("model", "")) usd_cost = usd_cost * provider_fee input_usd = input_usd * provider_fee output_usd = output_usd * provider_fee diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 89f283ed..64c03317 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -841,6 +841,7 @@ class BaseUpstreamProvider: new_session, max_cost_for_model, model_obj, + self.provider_fee, ) usage_finalized = True except Exception: @@ -1014,6 +1015,7 @@ class BaseUpstreamProvider: session, max_cost_for_model, model_obj, + self.provider_fee, ) usage_finalized = True except Exception as e: @@ -1158,6 +1160,7 @@ class BaseUpstreamProvider: session, deducted_max_cost, model_obj, + self.provider_fee, ) await session.refresh(key) @@ -1298,6 +1301,7 @@ class BaseUpstreamProvider: new_session, max_cost_for_model, model_obj, + self.provider_fee, ) usage_finalized = True except Exception: @@ -1428,6 +1432,7 @@ class BaseUpstreamProvider: session, max_cost_for_model, model_obj, + self.provider_fee, ) usage_finalized = True except Exception as e: @@ -1597,6 +1602,7 @@ class BaseUpstreamProvider: session, deducted_max_cost, model_obj, + self.provider_fee, ) await session.refresh(key) @@ -1797,6 +1803,7 @@ class BaseUpstreamProvider: new_session, max_cost_for_model, model_obj, + self.provider_fee, ) usage_finalized = True return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() @@ -1949,6 +1956,7 @@ class BaseUpstreamProvider: new_session, max_cost_for_model, model_obj, + self.provider_fee, ) self.inject_cost_metadata( @@ -2022,6 +2030,7 @@ class BaseUpstreamProvider: session, deducted_max_cost, model_obj, + self.provider_fee, ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2129,6 +2138,7 @@ class BaseUpstreamProvider: session, max_cost_for_model, model_obj, + self.provider_fee, ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2275,6 +2285,7 @@ class BaseUpstreamProvider: new_session, max_cost_for_model, model_obj, + self.provider_fee, ) usage_finalized = True return ( @@ -2349,6 +2360,7 @@ class BaseUpstreamProvider: new_session, max_cost_for_model, model_obj, + self.provider_fee, ) self.inject_cost_metadata( combined_data, cost_data, fresh_key @@ -3265,6 +3277,7 @@ class BaseUpstreamProvider: response_data, max_cost_for_model, model_obj, + self.provider_fee, ): case MaxCostData() as cost: logger.debug( diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index d7f919a7..ab98c270 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -502,7 +502,12 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: captured_cost_call: dict[str, Any] = {} async def fake_adjust( - fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None + fresh_key: Any, + combined_data: Any, + sess: Any, + max_cost: int, + model_obj: Any = None, + provider_fee: Any = None, ) -> dict: captured_cost_call["combined_data"] = combined_data captured_cost_call["max_cost"] = max_cost @@ -591,7 +596,12 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: captured: dict[str, Any] = {} async def fake_adjust( - fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None + fresh_key: Any, + combined_data: Any, + sess: Any, + max_cost: int, + model_obj: Any = None, + provider_fee: Any = None, ) -> dict: captured["combined_data"] = combined_data return fake_cost diff --git a/tests/unit/test_settlement_identity.py b/tests/unit/test_settlement_identity.py index 67eb00c2..6febb32d 100644 --- a/tests/unit/test_settlement_identity.py +++ b/tests/unit/test_settlement_identity.py @@ -88,6 +88,45 @@ async def test_string_fallback_still_prices_without_model_obj() -> None: assert result.total_msats == 2_000 +@pytest.mark.asyncio +async def test_usd_cost_path_applies_given_provider_fee() -> None: + """The USD-cost path bills the serving provider's fee when supplied.""" + from unittest.mock import Mock + + response = dict(RESPONSE) + response["usage"] = dict(RESPONSE["usage"], cost=0.001) # type: ignore[arg-type] + + best_ranked = Mock(provider_fee=1.0) + with patch( + "routstr.proxy.get_provider_for_model", return_value=[best_ranked] + ): + result = await calculate_cost( + response, max_cost=100_000, model_obj=SERVED, provider_fee=1.5 + ) + + assert isinstance(result, CostData) + # 0.001 USD * fee 1.5 / 0.0005 USD-per-sat = 3 sats = 3000 msats. + assert result.total_msats == 3_000 + + +@pytest.mark.asyncio +async def test_usd_cost_path_falls_back_to_best_ranked_fee() -> None: + """Without a supplied fee, the alias-map provider lookup still applies.""" + from unittest.mock import Mock + + response = dict(RESPONSE) + response["usage"] = dict(RESPONSE["usage"], cost=0.001) # type: ignore[arg-type] + + best_ranked = Mock(provider_fee=2.0) + with patch( + "routstr.proxy.get_provider_for_model", return_value=[best_ranked] + ): + result = await calculate_cost(response, max_cost=100_000) + + assert isinstance(result, CostData) + assert result.total_msats == 4_000 + + @pytest.mark.asyncio async def test_x_cashu_cost_uses_served_model_not_upstream_echo() -> None: """``get_x_cashu_cost`` bills the routed model, not the raw model echo.