diff --git a/router/payment/helpers.py b/router/payment/helpers.py index 61b8120a..002b4093 100644 --- a/router/payment/helpers.py +++ b/router/payment/helpers.py @@ -1,5 +1,6 @@ import json import os +from typing import Optional from fastapi import HTTPException, Response @@ -148,7 +149,9 @@ def get_max_cost_for_model(model: str) -> int: return COST_PER_REQUEST -def create_error_response(error_type: str, message: str, status_code: int) -> Response: +def create_error_response( + error_type: str, message: str, status_code: int, token: Optional[str] = None +) -> Response: """Create a standardized error response.""" logger.info( "Creating error response", @@ -159,6 +162,9 @@ def create_error_response(error_type: str, message: str, status_code: int) -> Re }, ) + response_headers = {} + if token: + response_headers["X-Cashu"] = token return Response( content=json.dumps( { @@ -171,6 +177,7 @@ def create_error_response(error_type: str, message: str, status_code: int) -> Re ), status_code=status_code, media_type="application/json", + headers=dict(response_headers), ) diff --git a/router/payment/x_cashu.py b/router/payment/x_cashu.py index 591ae837..afbc0d7e 100644 --- a/router/payment/x_cashu.py +++ b/router/payment/x_cashu.py @@ -8,12 +8,7 @@ from fastapi.responses import Response, StreamingResponse from ..core import get_logger from ..wallet import CurrencyUnit, recieve_token, send_token -from .cost_caculation import ( - CostData, - CostDataError, - MaxCostData, - calculate_cost, -) +from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost from .helpers import ( UPSTREAM_BASE_URL, create_error_response, @@ -68,21 +63,30 @@ async def x_cashu_handler( "token_already_spent", "The provided CASHU token has already been spent", 400, + x_cashu_token, ) - elif "invalid token" in error_message.lower(): + + if "invalid token" in error_message.lower(): return create_error_response( - "invalid_token", "The provided CASHU token is invalid", 400 + "invalid_token", + "The provided CASHU token is invalid", + 400, + x_cashu_token, ) - elif "mint error" in error_message.lower(): + + if "mint error" in error_message.lower(): return create_error_response( - "mint_error", f"CASHU mint error: {error_message}", 422 - ) - else: - # Generic error for other cases - return create_error_response( - "cashu_error", f"CASHU token processing failed: {error_message}", 400 + "mint_error", f"CASHU mint error: {error_message}", 422, x_cashu_token ) + # Generic error for other cases + return create_error_response( + "cashu_error", + f"CASHU token processing failed: {error_message}", + 400, + x_cashu_token, + ) + async def forward_to_upstream( request: Request, path: str, headers: dict, amount: int, unit: CurrencyUnit