diff --git a/router/cashu.py b/router/cashu.py index 8c8a8749..f7dd857d 100644 --- a/router/cashu.py +++ b/router/cashu.py @@ -203,6 +203,14 @@ async def refund_balance(amount_msats: int, key: ApiKey, session: AsyncSession) ) +async def x_cashu_refund(key: ApiKey, session: AsyncSession) -> str: + async with WALLET_LOCK: + refund_token = await WALLET.send(key.balance) + await session.delete(key) + await session.commit() + return refund_token + + async def redeem(cashu_token: str, lnurl: str) -> int: async with WALLET_LOCK: amount_sats, _ = await WALLET.redeem(cashu_token) diff --git a/router/proxy.py b/router/proxy.py index 26438c32..a9dbcd35 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -1,15 +1,16 @@ import json import os import re +from hashlib import sha256 from typing import AsyncGenerator import httpx -from fastapi import APIRouter, BackgroundTasks, Depends, Request +from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request from fastapi.responses import Response, StreamingResponse from .auth import adjust_payment_for_tokens, pay_for_request, validate_bearer_key -from .cashu import pay_out -from .db import AsyncSession, create_session, get_session +from .cashu import x_cashu_refund +from .db import ApiKey, AsyncSession, create_session, get_session UPSTREAM_BASE_URL = os.environ["UPSTREAM_BASE_URL"] UPSTREAM_API_KEY = os.environ.get("UPSTREAM_API_KEY", "") @@ -17,85 +18,68 @@ UPSTREAM_API_KEY = os.environ.get("UPSTREAM_API_KEY", "") proxy_router = APIRouter() -@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) -async def proxy( - request: Request, path: str, session: AsyncSession = Depends(get_session) -) -> Response | StreamingResponse: - auth = request.headers.get("Authorization", "") - bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else "" - refund_address = request.headers.get("Refund-LNURL", None) - key_expiry_time = request.headers.get("Key-Expiry-Time", None) +async def validate_and_read_json_body( + request: Request, path: str +) -> tuple[bytes | None, Response | None]: + """Validate and read JSON body for requests that require it. - # Validate key_expiry_time header - if key_expiry_time: - try: - key_expiry_time = int(key_expiry_time) # type: ignore - except ValueError: - return Response( - content="Invalid Key-Expiry-Time: must be a valid Unix timestamp", - status_code=400, - ) - if not refund_address: - return Response( - content="Error: Refund-LNURL header required when using Key-Expiry-Time", - status_code=400, - ) - else: - key_expiry_time = None + Returns: + tuple of (request_body, error_response) + - If successful: (body_bytes, None) + - If error: (None, error_response) + """ + if request.method not in ["POST", "PUT", "PATCH"] or not path.endswith( + "chat/completions" + ): + return None, None - key = await validate_bearer_key( - bearer_key, - session, - refund_address, - key_expiry_time, # type: ignore - ) - - # Pre-validate JSON for requests that require it - request_body = None - if request.method in ["POST", "PUT", "PATCH"] and path.endswith("chat/completions"): - try: - request_body = await request.body() - # Try to parse JSON to validate it - if request_body: - json.loads(request_body) - except json.JSONDecodeError as e: - return Response( - content=json.dumps( - { - "error": { - "message": f"Invalid JSON in request body: {str(e)}", - "type": "invalid_request_error", - "code": "invalid_json", - } + try: + request_body = await request.body() + # Try to parse JSON to validate it + if request_body: + json.loads(request_body) + return request_body, None + except json.JSONDecodeError as e: + return None, Response( + content=json.dumps( + { + "error": { + "message": f"Invalid JSON in request body: {str(e)}", + "type": "invalid_request_error", + "code": "invalid_json", } - ), - status_code=400, - media_type="application/json", - ) - except Exception: - return Response( - content=json.dumps( - { - "error": { - "message": "Error reading request body", - "type": "invalid_request_error", - "code": "request_error", - } + } + ), + status_code=400, + media_type="application/json", + ) + except Exception: + return None, Response( + content=json.dumps( + { + "error": { + "message": "Error reading request body", + "type": "invalid_request_error", + "code": "request_error", } - ), - status_code=400, - media_type="application/json", - ) + } + ), + status_code=400, + media_type="application/json", + ) - await pay_for_request(key, session, request, request_body) - # Prepare headers, removing sensitive/problematic ones - headers = dict(request.headers) +def prepare_upstream_headers(request_headers: dict) -> dict: + """Prepare headers for upstream request, removing sensitive/problematic ones.""" + headers = dict(request_headers) + # Remove headers that shouldn't be forwarded headers.pop("host", None) headers.pop("content-length", None) headers.pop("refund-lnurl", None) headers.pop("key-expiry-time", None) + headers.pop("x-cashu", None) + # Handle authorization if UPSTREAM_API_KEY: headers["Authorization"] = f"Bearer {UPSTREAM_API_KEY}" headers.pop("authorization", None) @@ -103,6 +87,134 @@ async def proxy( headers.pop("Authorization", None) headers.pop("authorization", None) + return headers + + +def create_error_response(error_type: str, message: str, status_code: int) -> Response: + """Create a standardized error response.""" + return Response( + content=json.dumps( + { + "error": { + "message": message, + "type": error_type, + "code": status_code, + } + } + ), + status_code=status_code, + media_type="application/json", + ) + + +async def handle_streaming_chat_completion( + response: httpx.Response, key: ApiKey, session: AsyncSession +) -> StreamingResponse: + """Handle streaming chat completion responses with token-based pricing.""" + + async def stream_with_cost() -> AsyncGenerator[bytes, None]: + # Store all chunks to analyze + stored_chunks = [] + + async for chunk in response.aiter_bytes(): + # Store chunk for later analysis + stored_chunks.append(chunk) + + # Pass through each chunk to client + yield chunk + + # Process stored chunks to find usage data + # Start from the end and work backwards + for i in range(len(stored_chunks) - 1, -1, -1): + chunk = stored_chunks[i] + if not chunk or chunk == b"": + continue + + try: + # Split by "data: " to get individual SSE events + events = re.split(b"data: ", chunk) + for event_data in events: + if ( + not event_data + or event_data.strip() == b"[DONE]" + or event_data.strip() == b"" + ): + continue + + try: + data = json.loads(event_data) + if ( + "usage" in data + and data["usage"] is not None + and isinstance(data["usage"], dict) + ): + # Found usage data, calculate cost + # Create a new session for this operation + async with create_session() as new_session: + # Re-fetch the key in the new session + fresh_key = await new_session.get( + key.__class__, key.hashed_key + ) + if fresh_key: + cost_data = await adjust_payment_for_tokens( + fresh_key, data, new_session + ) + # Format as SSE and yield + cost_json = json.dumps({"cost": cost_data}) + yield f"data: {cost_json}\n\n".encode() + break + except json.JSONDecodeError: + continue + + except Exception as e: + print(f"Error processing streaming response for cost: {e}") + + return StreamingResponse( + stream_with_cost(), + status_code=response.status_code, + headers=dict(response.headers), + ) + + +async def handle_non_streaming_chat_completion( + response: httpx.Response, key: ApiKey, session: AsyncSession +) -> Response: + """Handle non-streaming chat completion responses with token-based pricing.""" + try: + content = await response.aread() + response_json = json.loads(content) + cost_data = await adjust_payment_for_tokens(key, response_json, session) + response_json["cost"] = cost_data + + response_headers = dict(response.headers) + + # Remove Transfer-Encoding header to avoid conflict with Content-Length header in common nginx setups + if "transfer-encoding" in response_headers: + del response_headers["transfer-encoding"] + + return Response( + content=json.dumps(response_json).encode(), + status_code=response.status_code, + headers=response_headers, + media_type="application/json", + ) + except json.JSONDecodeError as e: + print(f"Failed to parse JSON from upstream response: {e}") + raise + except Exception as e: + print(f"Error adjusting payment for tokens: {e}") + raise + + +async def forward_to_upstream( + request: Request, + path: str, + headers: dict, + request_body: bytes | None, + key: ApiKey, + session: AsyncSession, +) -> Response | StreamingResponse: + """Forward request to upstream and handle the response.""" if path.startswith("v1/"): path = path.replace("v1/", "") @@ -145,103 +257,19 @@ async def proxy( if is_streaming and response.status_code == 200: # Process streaming response and extract cost from the last chunk - async def stream_with_cost() -> AsyncGenerator[bytes, None]: - # Store all chunks to analyze - stored_chunks = [] - - async for chunk in response.aiter_bytes(): - # Store chunk for later analysis - stored_chunks.append(chunk) - - # Pass through each chunk to client - yield chunk - - # Process stored chunks to find usage data - # Start from the end and work backwards - for i in range(len(stored_chunks) - 1, -1, -1): - chunk = stored_chunks[i] - if not chunk or chunk == b"": - continue - - try: - # Split by "data: " to get individual SSE events - events = re.split(b"data: ", chunk) - for event_data in events: - if ( - not event_data - or event_data.strip() == b"[DONE]" - or event_data.strip() == b"" - ): - continue - - try: - data = json.loads(event_data) - if ( - "usage" in data - and data["usage"] is not None - and isinstance(data["usage"], dict) - ): - # Found usage data, calculate cost - # Create a new session for this operation - async with create_session() as new_session: - # Re-fetch the key in the new session - fresh_key = await new_session.get( - key.__class__, key.hashed_key - ) - if fresh_key: - cost_data = ( - await adjust_payment_for_tokens( - fresh_key, data, new_session - ) - ) - # Format as SSE and yield - cost_json = json.dumps( - {"cost": cost_data} - ) - yield f"data: {cost_json}\n\n".encode() - break - except json.JSONDecodeError: - continue - - except Exception as e: - print(f"Error processing streaming response for cost: {e}") - + result = await handle_streaming_chat_completion(response, key, session) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) background_tasks.add_task(client.aclose) - return StreamingResponse( - stream_with_cost(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, - ) + result.background = background_tasks + return result elif response.status_code == 200 and "application/json" in content_type: # Handle non-streaming response try: - content = await response.aread() - response_json = json.loads(content) - cost_data = await adjust_payment_for_tokens( - key, response_json, session + return await handle_non_streaming_chat_completion( + response, key, session ) - response_json["cost"] = cost_data - - response_headers = dict(response.headers) - - # Remove Transfer-Encoding header to avoid conflict with Content-Length header in common nginx setups - if "transfer-encoding" in response_headers: - del response_headers["transfer-encoding"] - - return Response( - content=json.dumps(response_json).encode(), - status_code=response.status_code, - headers=response_headers, - media_type="application/json", - ) - except json.JSONDecodeError as e: - print(f"Failed to parse JSON from upstream response: {e}") - except Exception as e: - print(f"Error adjusting payment for tokens: {e}") finally: await response.aclose() await client.aclose() @@ -250,7 +278,6 @@ async def proxy( background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) background_tasks.add_task(client.aclose) - background_tasks.add_task(pay_out) return StreamingResponse( response.aiter_bytes(), @@ -279,19 +306,8 @@ async def proxy( else: error_message = f"Error connecting to upstream service: {error_type}" - return Response( - content=json.dumps( - { - "error": { - "message": error_message, - "type": "upstream_error", - "code": 502, - } - } - ), - status_code=502, - media_type="application/json", - ) + return create_error_response("upstream_error", error_message, 502) + except Exception as exc: await client.aclose() import traceback @@ -303,16 +319,78 @@ async def proxy( f"path={path}, query_params={dict(request.query_params)}\n" f"Traceback:\n{tb}" ) - return Response( - content=json.dumps( - { - "error": { - "message": "An unexpected server error occurred", - "type": "internal_error", - "code": 500, - } - } - ), - status_code=500, - media_type="application/json", + return create_error_response( + "internal_error", "An unexpected server error occurred", 500 ) + + +@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) +async def proxy( + request: Request, path: str, session: AsyncSession = Depends(get_session) +) -> Response | StreamingResponse: + # Check for X-Cashu header first + if x_cashu := request.headers.get("X-Cashu", None): + key = await validate_bearer_key( + sha256(x_cashu.encode()).hexdigest(), session, "X-CASHU" + ) + + if auth := request.headers.get("Authorization", None): + key = await get_bearer_token_key(request, path, session, auth) + + else: + raise HTTPException(status_code=401, detail="Unauthorized") + + request_body, error_response = await validate_and_read_json_body(request, path) + if error_response: + return error_response + + await pay_for_request(key, session, request, request_body) + + # Prepare headers for upstream + headers = prepare_upstream_headers(dict(request.headers)) + + # Forward to upstream and handle response + response = await forward_to_upstream( + request, path, headers, request_body, key, session + ) + + if key.refund_address == "X-CASHU": + refund_token = await x_cashu_refund(key, session) + response.headers["X-Cashu-Refund"] = refund_token + + return response + + +async def get_bearer_token_key( + request: Request, path: str, session: AsyncSession, auth: str +) -> ApiKey: + """Handle bearer token authentication proxy requests.""" + # Handle regular bearer token authentication + auth = request.headers.get("Authorization", "") + bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else "" + refund_address = request.headers.get("Refund-LNURL", None) + key_expiry_time = request.headers.get("Key-Expiry-Time", None) + + # Validate key_expiry_time header + if key_expiry_time: + try: + key_expiry_time = int(key_expiry_time) # type: ignore + except ValueError: + raise HTTPException( + status_code=400, + detail="Invalid Key-Expiry-Time: must be a valid Unix timestamp", + ) + if not refund_address: + raise HTTPException( + status_code=400, + detail="Error: Refund-LNURL header required when using Key-Expiry-Time", + ) + else: + key_expiry_time = None + + return await validate_bearer_key( + bearer_key, + session, + refund_address, + key_expiry_time, # type: ignore + )