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,
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)",
+12 -1
View File
@@ -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,6 +388,7 @@ def _get_pricing_rates(
if input_usd <= 0 or output_usd <= 0:
raise ValueError(f"Incomplete LiteLLM pricing for model: {pricing_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
@@ -441,8 +450,10 @@ 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."""
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
+13
View File
@@ -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(
+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] = {}
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
+39
View File
@@ -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.