From ffae549d330a4a5ab15fb8bcb08e3309ffa07b7f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 3 Sep 2026 01:05:34 +0200 Subject: [PATCH] clean up xcashu path --- routstr/upstream/base.py | 39 ++++++++++++++++---- tests/unit/test_messages_litellm_dispatch.py | 32 +++++++++++----- 2 files changed, 54 insertions(+), 17 deletions(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 420832a1..39005698 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -255,6 +255,13 @@ def _openai_completion_path(path: str) -> str | None: return "completions" if canonical.endswith("/completions") else None +def _x_cashu_path_has_settlement_handler(path: str) -> bool: + canonical = path.rstrip("/") + return _openai_completion_path(canonical) is not None or canonical.endswith( + ("embeddings", "messages", "messages/count_tokens") + ) + + class TopupData(BaseModel): """Universal top-up data schema for Lightning Network invoices.""" @@ -4245,7 +4252,15 @@ class BaseUpstreamProvider: path.endswith("messages/count_tokens") and not self.supports_anthropic_messages ): - return count_tokens_locally(request_body, model_obj) + result = count_tokens_locally(request_body, model_obj) + refund_token = await self.send_refund( + amount, + unit, + mint, + request_id=getattr(request.state, "request_id", None), + ) + result.headers["X-Cashu"] = refund_token + return result if ( path.endswith("messages") @@ -4364,12 +4379,7 @@ class BaseUpstreamProvider: error_response.headers["X-Cashu"] = refund_token return error_response - if ( - completion_path is not None - or path.endswith("embeddings") - or path.endswith("messages") - or path.endswith("messages/count_tokens") - ): + if _x_cashu_path_has_settlement_handler(path): logger.debug( "Processing completion/embeddings/messages response", extra={"path": path, "amount": amount, "unit": unit}, @@ -5157,6 +5167,21 @@ class BaseUpstreamProvider: }, ) + # Reject before redemption so the client keeps its token. + if not _x_cashu_path_has_settlement_handler(path): + logger.warning( + "Rejecting X-Cashu request for unsupported endpoint", + extra={"path": path, "method": request.method}, + ) + return create_error_response( + "invalid_request_error", + "X-Cashu payment is not supported on this endpoint; use bearer " + "(deposit) authentication instead. The token was not redeemed.", + 400, + request=request, + code="x_cashu_unsupported_endpoint", + ) + redeemed = False try: headers = dict(request.headers) diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index ec93d555..ca5a83c5 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -1134,18 +1134,30 @@ async def test_forward_x_cashu_request_handles_count_tokens_locally() -> None: "prepare_request_body", side_effect=AssertionError("upstream should not be called"), ): - response = await provider.forward_x_cashu_request( - request=request, - path="v1/messages/count_tokens", - headers={}, - amount=5_000, - unit="sat", - max_cost_for_model=10_000, - model_obj=model, - mint="https://mint", - ) + with patch.object( + provider, + "send_refund", + new=AsyncMock(return_value="refund-token"), + ) as send_refund: + response = await provider.forward_x_cashu_request( + request=request, + path="v1/messages/count_tokens", + headers={}, + amount=5_000, + unit="sat", + max_cost_for_model=10_000, + model_obj=model, + mint="https://mint", + ) + send_refund.assert_awaited_once_with( + 5_000, + "sat", + "https://mint", + request_id="req-test", + ) assert response.status_code == 200 + assert response.headers["X-Cashu"] == "refund-token" body = response.body if isinstance(response.body, bytes) else bytes(response.body) payload = json.loads(body.decode()) assert "input_tokens" in payload