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