clean up xcashu path

This commit is contained in:
9qeklajc
2026-09-03 01:05:34 +02:00
parent 0759e024aa
commit ffae549d33
2 changed files with 54 additions and 17 deletions
+32 -7
View File
@@ -255,6 +255,13 @@ def _openai_completion_path(path: str) -> str | None:
return "completions" if canonical.endswith("/completions") else 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): class TopupData(BaseModel):
"""Universal top-up data schema for Lightning Network invoices.""" """Universal top-up data schema for Lightning Network invoices."""
@@ -4245,7 +4252,15 @@ class BaseUpstreamProvider:
path.endswith("messages/count_tokens") path.endswith("messages/count_tokens")
and not self.supports_anthropic_messages 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 ( if (
path.endswith("messages") path.endswith("messages")
@@ -4364,12 +4379,7 @@ class BaseUpstreamProvider:
error_response.headers["X-Cashu"] = refund_token error_response.headers["X-Cashu"] = refund_token
return error_response return error_response
if ( if _x_cashu_path_has_settlement_handler(path):
completion_path is not None
or path.endswith("embeddings")
or path.endswith("messages")
or path.endswith("messages/count_tokens")
):
logger.debug( logger.debug(
"Processing completion/embeddings/messages response", "Processing completion/embeddings/messages response",
extra={"path": path, "amount": amount, "unit": unit}, 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 redeemed = False
try: try:
headers = dict(request.headers) headers = dict(request.headers)
+22 -10
View File
@@ -1134,18 +1134,30 @@ async def test_forward_x_cashu_request_handles_count_tokens_locally() -> None:
"prepare_request_body", "prepare_request_body",
side_effect=AssertionError("upstream should not be called"), side_effect=AssertionError("upstream should not be called"),
): ):
response = await provider.forward_x_cashu_request( with patch.object(
request=request, provider,
path="v1/messages/count_tokens", "send_refund",
headers={}, new=AsyncMock(return_value="refund-token"),
amount=5_000, ) as send_refund:
unit="sat", response = await provider.forward_x_cashu_request(
max_cost_for_model=10_000, request=request,
model_obj=model, path="v1/messages/count_tokens",
mint="https://mint", 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.status_code == 200
assert response.headers["X-Cashu"] == "refund-token"
body = response.body if isinstance(response.body, bytes) else bytes(response.body) body = response.body if isinstance(response.body, bytes) else bytes(response.body)
payload = json.loads(body.decode()) payload = json.loads(body.decode())
assert "input_tokens" in payload assert "input_tokens" in payload