fix: bill the serving provider's fee on the USD-cost path

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 <noreply@anthropic.com>
This commit is contained in:
Jeroen Ubbink
2026-07-22 15:09:41 +02:00
co-authored by Claude Fable 5
parent b76fa17f81
commit 0aebfc6dbe
5 changed files with 82 additions and 6 deletions
+4 -1
View File
@@ -775,6 +775,7 @@ async def adjust_payment_for_tokens(
session: AsyncSession, session: AsyncSession,
deducted_max_cost: int, deducted_max_cost: int,
model_obj: "Model | None" = None, model_obj: "Model | None" = None,
provider_fee: float | None = None,
) -> dict: ) -> dict:
""" """
Adjusts the payment based on token usage in the response. 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}, 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: case MaxCostData() as cost:
logger.debug( logger.debug(
"Using max cost data (no token adjustment)", "Using max cost data (no token adjustment)",
+14 -3
View File
@@ -71,6 +71,7 @@ async def calculate_cost(
response_data: dict, response_data: dict,
max_cost: int, max_cost: int,
model_obj: "Model | None" = None, model_obj: "Model | None" = None,
provider_fee: float | None = None,
) -> CostData | MaxCostData | CostDataError: ) -> CostData | MaxCostData | CostDataError:
"""Calculate the cost of an API request based on token usage. """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 its pricing is billed directly; without it, pricing is re-derived
from the response's model string via the alias map, which resolves from the response's model string via the alias map, which resolves
to the best-ranked candidate — not necessarily the serving one. 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: Returns:
Cost data or error information Cost data or error information
@@ -186,6 +192,7 @@ async def calculate_cost(
cache_creation_tokens, cache_creation_tokens,
output_tokens, output_tokens,
response_data, response_data,
provider_fee,
) )
except Exception as e: except Exception as e:
logger.warning( logger.warning(
@@ -199,7 +206,7 @@ async def calculate_cost(
# Fall back to token-based pricing # Fall back to token-based pricing
try: 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: except ValueError as e:
return CostDataError(message=str(e), code="pricing_error") 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( def _get_pricing_rates(
response_data: dict, response_data: dict,
model_obj: "Model | None" = None, model_obj: "Model | None" = None,
provider_fee: float | None = None,
) -> tuple[float, float, float, float] | None: ) -> tuple[float, float, float, float] | None:
"""Get configured rates, falling back to LiteLLM's model cost map. """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: if input_usd <= 0 or output_usd <= 0:
raise ValueError(f"Incomplete LiteLLM pricing for model: {pricing_model}") 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() usd_per_sat = sats_usd_price()
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat 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 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, cache_creation_tokens: int,
output_tokens: int, output_tokens: int,
response_data: dict, response_data: dict,
provider_fee: float | None = None,
) -> CostData: ) -> CostData:
"""Calculate cost from USD figures, deriving input/output split from tokens.""" """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 usd_cost = usd_cost * provider_fee
input_usd = input_usd * provider_fee input_usd = input_usd * provider_fee
output_usd = output_usd * provider_fee output_usd = output_usd * provider_fee
+13
View File
@@ -841,6 +841,7 @@ class BaseUpstreamProvider:
new_session, new_session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
usage_finalized = True usage_finalized = True
except Exception: except Exception:
@@ -1014,6 +1015,7 @@ class BaseUpstreamProvider:
session, session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
usage_finalized = True usage_finalized = True
except Exception as e: except Exception as e:
@@ -1158,6 +1160,7 @@ class BaseUpstreamProvider:
session, session,
deducted_max_cost, deducted_max_cost,
model_obj, model_obj,
self.provider_fee,
) )
await session.refresh(key) await session.refresh(key)
@@ -1298,6 +1301,7 @@ class BaseUpstreamProvider:
new_session, new_session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
usage_finalized = True usage_finalized = True
except Exception: except Exception:
@@ -1428,6 +1432,7 @@ class BaseUpstreamProvider:
session, session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
usage_finalized = True usage_finalized = True
except Exception as e: except Exception as e:
@@ -1597,6 +1602,7 @@ class BaseUpstreamProvider:
session, session,
deducted_max_cost, deducted_max_cost,
model_obj, model_obj,
self.provider_fee,
) )
await session.refresh(key) await session.refresh(key)
@@ -1797,6 +1803,7 @@ class BaseUpstreamProvider:
new_session, new_session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
usage_finalized = True usage_finalized = True
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
@@ -1949,6 +1956,7 @@ class BaseUpstreamProvider:
new_session, new_session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
self.inject_cost_metadata( self.inject_cost_metadata(
@@ -2022,6 +2030,7 @@ class BaseUpstreamProvider:
session, session,
deducted_max_cost, deducted_max_cost,
model_obj, model_obj,
self.provider_fee,
) )
self.inject_cost_metadata(response_json, cost_data, key) self.inject_cost_metadata(response_json, cost_data, key)
@@ -2129,6 +2138,7 @@ class BaseUpstreamProvider:
session, session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
self.inject_cost_metadata(response_json, cost_data, key) self.inject_cost_metadata(response_json, cost_data, key)
@@ -2275,6 +2285,7 @@ class BaseUpstreamProvider:
new_session, new_session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
usage_finalized = True usage_finalized = True
return ( return (
@@ -2349,6 +2360,7 @@ class BaseUpstreamProvider:
new_session, new_session,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
) )
self.inject_cost_metadata( self.inject_cost_metadata(
combined_data, cost_data, fresh_key combined_data, cost_data, fresh_key
@@ -3265,6 +3277,7 @@ class BaseUpstreamProvider:
response_data, response_data,
max_cost_for_model, max_cost_for_model,
model_obj, model_obj,
self.provider_fee,
): ):
case MaxCostData() as cost: case MaxCostData() as cost:
logger.debug( logger.debug(
+12 -2
View File
@@ -502,7 +502,12 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
captured_cost_call: dict[str, Any] = {} captured_cost_call: dict[str, Any] = {}
async def fake_adjust( 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: ) -> dict:
captured_cost_call["combined_data"] = combined_data captured_cost_call["combined_data"] = combined_data
captured_cost_call["max_cost"] = max_cost 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] = {} captured: dict[str, Any] = {}
async def fake_adjust( 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: ) -> dict:
captured["combined_data"] = combined_data captured["combined_data"] = combined_data
return fake_cost return fake_cost
+39
View File
@@ -88,6 +88,45 @@ async def test_string_fallback_still_prices_without_model_obj() -> None:
assert result.total_msats == 2_000 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 @pytest.mark.asyncio
async def test_x_cashu_cost_uses_served_model_not_upstream_echo() -> None: 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. """``get_x_cashu_cost`` bills the routed model, not the raw model echo.