From b76fa17f81c1f008c20986d7c47a2d1e24abd72b Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Fri, 17 Jul 2026 15:57:02 +0200 Subject: [PATCH] fix: thread the served model into x-cashu settlement pricing MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- routstr/upstream/base.py | 41 +++++++++++++++++++++----- tests/unit/test_settlement_identity.py | 25 ++++++++++++++++ 2 files changed, 59 insertions(+), 7 deletions(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 35996ac8..89f283ed 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -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 diff --git a/tests/unit/test_settlement_identity.py b/tests/unit/test_settlement_identity.py index b1300df8..67eb00c2 100644 --- a/tests/unit/test_settlement_identity.py +++ b/tests/unit/test_settlement_identity.py @@ -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()