diff --git a/routstr/auth.py b/routstr/auth.py index 78b503c4..4c891c96 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -272,10 +272,10 @@ async def validate_bearer_key( ) -async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int: +async def pay_for_request( + key: ApiKey, cost_per_request: int, session: AsyncSession +) -> int: """Process payment for a request.""" - model = body["model"] - cost_per_request = get_max_cost_for_model(model=model) logger.info( "Processing payment for request", @@ -283,7 +283,6 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int "key_hash": key.hashed_key[:8] + "...", "current_balance": key.balance, "required_cost": cost_per_request, - "model": model, "sufficient_balance": key.balance >= cost_per_request, }, ) @@ -297,7 +296,6 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int "reserved_balance": key.reserved_balance, "required": cost_per_request, "shortfall": cost_per_request - key.total_balance, - "model": model, }, ) @@ -366,7 +364,6 @@ async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> int "new_balance": key.balance, "total_spent": key.total_spent, "total_requests": key.total_requests, - "model": model, }, ) diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 4404fd0b..83bf7511 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -19,30 +19,6 @@ if not UPSTREAM_BASE_URL: raise ValueError("Please set the UPSTREAM_BASE_URL environment variable") -def get_cost_per_request(model: str | None = None) -> int: - """Get the cost per request for a given model.""" - logger.debug( - "Calculating cost per request", - extra={ - "model": model, - "model_based_pricing": MODEL_BASED_PRICING, - "has_models": bool(MODELS), - }, - ) - - if MODEL_BASED_PRICING and MODELS and model: - cost = get_max_cost_for_model(model=model) - logger.debug( - "Using model-based cost", extra={"model": model, "cost_msats": cost} - ) - return cost - - logger.debug( - "Using default cost per request", extra={"cost_msats": COST_PER_REQUEST} - ) - return COST_PER_REQUEST - - def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None: if x_cashu := headers.get("x-cashu", None): cashu_token = x_cashu diff --git a/routstr/payment/x_cashu.py b/routstr/payment/x_cashu.py index a51df5e1..551dce4a 100644 --- a/routstr/payment/x_cashu.py +++ b/routstr/payment/x_cashu.py @@ -9,18 +9,13 @@ from fastapi.responses import Response, StreamingResponse from ..core import get_logger from ..wallet import recieve_token, send_token from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost -from .helpers import ( - UPSTREAM_BASE_URL, - create_error_response, - get_max_cost_for_model, - prepare_upstream_headers, -) +from .helpers import UPSTREAM_BASE_URL, create_error_response, prepare_upstream_headers logger = get_logger(__name__) async def x_cashu_handler( - request: Request, x_cashu_token: str, path: str + request: Request, x_cashu_token: str, path: str, max_cost_for_model: int ) -> Response | StreamingResponse: """Handle X-Cashu token payment requests.""" logger.info( @@ -44,7 +39,9 @@ async def x_cashu_handler( extra={"amount": amount, "unit": unit, "path": path, "mint": mint}, ) - return await forward_to_upstream(request, path, headers, amount, unit) + return await forward_to_upstream( + request, path, headers, amount, unit, max_cost_for_model + ) except Exception as e: error_message = str(e) logger.error( @@ -96,7 +93,12 @@ async def x_cashu_handler( async def forward_to_upstream( - request: Request, path: str, headers: dict, amount: int, unit: str + request: Request, + path: str, + headers: dict, + amount: int, + unit: str, + max_cost_for_model: int, ) -> Response | StreamingResponse: """Forward request to upstream and handle the response.""" if path.startswith("v1/"): @@ -188,7 +190,9 @@ async def forward_to_upstream( extra={"path": path, "amount": amount, "unit": unit}, ) - result = await handle_x_cashu_chat_completion(response, amount, unit) + result = await handle_x_cashu_chat_completion( + response, amount, unit, max_cost_for_model + ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) result.background = background_tasks @@ -232,7 +236,7 @@ async def forward_to_upstream( async def handle_x_cashu_chat_completion( - response: httpx.Response, amount: int, unit: str + response: httpx.Response, amount: int, unit: str, max_cost_for_model: int ) -> StreamingResponse | Response: """Handle both streaming and non-streaming chat completion responses with token-based pricing.""" logger.debug( @@ -256,10 +260,12 @@ async def handle_x_cashu_chat_completion( ) if is_streaming: - return await handle_streaming_response(content_str, response, amount, unit) + return await handle_streaming_response( + content_str, response, amount, unit, max_cost_for_model + ) else: return await handle_non_streaming_response( - content_str, response, amount, unit + content_str, response, amount, unit, max_cost_for_model ) except Exception as e: @@ -281,7 +287,11 @@ async def handle_x_cashu_chat_completion( async def handle_streaming_response( - content_str: str, response: httpx.Response, amount: int, unit: str + content_str: str, + response: httpx.Response, + amount: int, + unit: str, + max_cost_for_model: int, ) -> StreamingResponse: """Handle Server-Sent Events (SSE) streaming response.""" logger.debug( @@ -335,7 +345,7 @@ async def handle_streaming_response( response_data = {"usage": usage_data, "model": model} try: - cost_data = await get_cost(response_data) + cost_data = await get_cost(response_data, max_cost_for_model) if cost_data: if unit == "msat": refund_amount = amount - cost_data.total_msats @@ -403,7 +413,11 @@ async def handle_streaming_response( async def handle_non_streaming_response( - content_str: str, response: httpx.Response, amount: int, unit: str + content_str: str, + response: httpx.Response, + amount: int, + unit: str, + max_cost_for_model: int, ) -> Response: """Handle regular JSON response.""" logger.debug( @@ -414,7 +428,7 @@ async def handle_non_streaming_response( try: response_json = json.loads(content_str) - cost_data = await get_cost(response_json) + cost_data = await get_cost(response_json, max_cost_for_model) if not cost_data: logger.error( @@ -520,21 +534,21 @@ async def handle_non_streaming_response( ) -async def get_cost(response_data: dict) -> MaxCostData | CostData | None: +async def get_cost( + response_data: dict, max_cost_for_model: int +) -> MaxCostData | CostData | None: """ Adjusts the payment based on token usage in the response. This is called after the initial payment and the upstream request is complete. Returns cost data to be included in the response. """ - model = response_data.get("model", "unknown") + model = response_data.get("model", None) logger.debug( "Calculating cost for response", extra={"model": model, "has_usage": "usage" in response_data}, ) - max_cost = get_max_cost_for_model(model=model) - - match calculate_cost(response_data, max_cost): + match calculate_cost(response_data, max_cost_for_model): case MaxCostData() as cost: logger.debug( "Using max cost pricing", diff --git a/routstr/proxy.py b/routstr/proxy.py index 2e0d615f..5ebe7c6f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -19,7 +19,7 @@ from .payment.helpers import ( UPSTREAM_BASE_URL, check_token_balance, create_error_response, - get_cost_per_request, + get_max_cost_for_model, prepare_upstream_headers, ) from .payment.x_cashu import x_cashu_handler @@ -501,9 +501,8 @@ async def proxy( media_type="application/json", ) - max_cost_for_model = get_cost_per_request( - model=request_body_dict.get("model", None) - ) + model = request_body_dict.get("model", "unknown") + max_cost_for_model = get_max_cost_for_model(model=model) check_token_balance(headers, request_body_dict, max_cost_for_model) # Handle authentication @@ -515,7 +514,7 @@ async def proxy( "token_preview": x_cashu[:20] + "..." if len(x_cashu) > 20 else x_cashu, }, ) - return await x_cashu_handler(request, x_cashu, path) + return await x_cashu_handler(request, x_cashu, path, max_cost_for_model) elif auth := headers.get("authorization", None): logger.debug( @@ -557,7 +556,7 @@ async def proxy( ) try: - await pay_for_request(key, session, request_body_dict) + await pay_for_request(key, max_cost_for_model, session) logger.info( "Payment processed successfully", extra={