mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
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:
co-authored by
Claude Fable 5
parent
b76fa17f81
commit
0aebfc6dbe
+4
-1
@@ -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)",
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
Reference in New Issue
Block a user