fix: thread the served model into x-cashu settlement pricing

X-Cashu handlers do not rewrite the upstream's echoed model string, so
get_x_cashu_cost previously priced whatever wire name the upstream
reported — the most collapse-prone alias lookup of all. The routed Model
is now threaded from forward_x_cashu_request through the chat and
Responses handler chains (and the litellm messages path) into
get_x_cashu_cost, so cost and refund are computed from the model that
actually served.

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 ed8ac914f9
commit b76fa17f81
2 changed files with 59 additions and 7 deletions
+34 -7
View File
@@ -2168,6 +2168,7 @@ class BaseUpstreamProvider:
requested_model,
mint,
request_id,
model_obj,
)
response_json = messages_dispatch.coerce_litellm_payload(result)
@@ -2175,7 +2176,9 @@ class BaseUpstreamProvider:
if requested_model and "model" in response_json:
response_json["model"] = requested_model
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data and "usage" in response_json and isinstance(
response_json["usage"], dict
@@ -2383,6 +2386,7 @@ class BaseUpstreamProvider:
requested_model: str | None,
mint: str | None,
request_id: str | None,
model_obj: Model | None = None,
) -> StreamingResponse:
"""Buffer a litellm stream end-to-end, compute cost, then replay.
@@ -2474,7 +2478,7 @@ class BaseUpstreamProvider:
}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model
response_data, max_cost_for_model, model_obj
)
if cost_data:
refund_amount = messages_dispatch.compute_refund(
@@ -3234,13 +3238,19 @@ class BaseUpstreamProvider:
)
async def get_x_cashu_cost(
self, response_data: dict, max_cost_for_model: int
self,
response_data: dict,
max_cost_for_model: int,
model_obj: Model | None = None,
) -> MaxCostData | CostData | None:
"""Calculate cost for X-Cashu payment based on response data.
Args:
response_data: Response data containing model and usage information
max_cost_for_model: Maximum cost for the model
model_obj: The model that actually served the request; billed
directly instead of re-deriving pricing from the upstream's
echoed model string
Returns:
Cost data object (MaxCostData or CostData) or None if calculation fails
@@ -3254,6 +3264,7 @@ class BaseUpstreamProvider:
match await calculate_cost(
response_data,
max_cost_for_model,
model_obj,
):
case MaxCostData() as cost:
logger.debug(
@@ -3398,6 +3409,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> StreamingResponse:
"""Handle streaming response for X-Cashu payment, calculating refund if needed.
@@ -3471,7 +3483,7 @@ class BaseUpstreamProvider:
response_data = {"usage": usage_data, "model": model}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model
response_data, max_cost_for_model, model_obj
)
if cost_data:
if unit == "msat":
@@ -3576,6 +3588,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> Response:
"""Handle non-streaming response for X-Cashu payment, calculating refund if needed.
@@ -3597,7 +3610,9 @@ class BaseUpstreamProvider:
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data and "usage" in response_json:
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
@@ -3726,6 +3741,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> StreamingResponse | Response:
"""Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming.
@@ -3769,6 +3785,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
)
else:
return await self.handle_x_cashu_non_streaming_response(
@@ -3779,6 +3796,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
)
except Exception as e:
@@ -3964,6 +3982,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=getattr(request.state, "request_id", None),
model_obj=model_obj,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -4256,6 +4275,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=getattr(request.state, "request_id", None),
model_obj=model_obj,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -4306,6 +4326,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> StreamingResponse | Response:
"""Handle Responses API completion response for X-Cashu payment.
@@ -4350,6 +4371,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
)
else:
return await self.handle_x_cashu_non_streaming_responses_response(
@@ -4360,6 +4382,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
)
except Exception as e:
@@ -4387,6 +4410,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> StreamingResponse:
"""Handle streaming Responses API response for X-Cashu payment.
@@ -4445,7 +4469,7 @@ class BaseUpstreamProvider:
response_data = {"usage": usage_data, "model": model}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model
response_data, max_cost_for_model, model_obj
)
if cost_data:
if unit == "msat":
@@ -4551,6 +4575,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> Response:
"""Handle non-streaming Responses API response for X-Cashu payment."""
logger.debug(
@@ -4561,7 +4586,9 @@ class BaseUpstreamProvider:
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data and "usage" in response_json:
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
+25
View File
@@ -86,3 +86,28 @@ async def test_string_fallback_still_prices_without_model_obj() -> None:
assert isinstance(result, CostData)
assert result.total_msats == 2_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.
X-Cashu handlers do not rewrite the upstream's echoed model string, so
without the routed model the settle would look up whatever wire name the
upstream reported. With ``model_obj`` given, the echo must be irrelevant.
"""
from routstr.upstream import GenericUpstreamProvider
provider = GenericUpstreamProvider("http://upstream.example", "key", 1.0)
response = dict(RESPONSE, model="totally-unknown-wire-name")
with patch(
"routstr.proxy.get_model_instance", return_value=WINNER
) as alias_lookup:
cost = await provider.get_x_cashu_cost(
response, max_cost_for_model=100_000, model_obj=SERVED
)
assert cost is not None
assert cost.total_msats == 10_000
alias_lookup.assert_not_called()