From 8beab0e28bfd264083a9b1e2795047316b3e19b5 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Mon, 6 Oct 2025 16:40:36 +0800 Subject: [PATCH] move functions into class to be overwritten by subclasses --- routstr/proxy.py | 14 +- routstr/upstream.py | 1618 +++++++++++++++++++++---------------------- 2 files changed, 780 insertions(+), 852 deletions(-) diff --git a/routstr/proxy.py b/routstr/proxy.py index abee68b7..394c92c0 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -13,12 +13,7 @@ from .payment.helpers import ( create_error_response, get_max_cost_for_model, ) -from .upstream import ( - UpstreamProvider, - handle_non_streaming_chat_completion, - handle_streaming_chat_completion, - map_upstream_error_response, -) +from .upstream import UpstreamProvider logger = get_logger(__name__) proxy_router = APIRouter() @@ -128,9 +123,7 @@ async def proxy( logger.debug("Processing unauthenticated GET request", extra={"path": path}) # TODO: why is this needed? can we remove it? headers = upstream.prepare_headers(dict(request.headers)) - return await upstream.forward_get_request( - request, path, headers, map_upstream_error_response - ) + return await upstream.forward_get_request(request, path, headers) # Only pay for request if we have request body data (for completions endpoints) if request_body_dict: @@ -179,9 +172,6 @@ async def proxy( key, max_cost_for_model, session, - map_upstream_error_response, - handle_streaming_chat_completion, - handle_non_streaming_chat_completion, ) if response.status_code != 200: diff --git a/routstr/upstream.py b/routstr/upstream.py index 23dca318..1f9dce6b 100644 --- a/routstr/upstream.py +++ b/routstr/upstream.py @@ -3,7 +3,7 @@ from __future__ import annotations import json import re import traceback -from collections.abc import AsyncGenerator, Awaitable, Callable +from collections.abc import AsyncGenerator from typing import TYPE_CHECKING, Mapping import httpx @@ -17,413 +17,11 @@ from fastapi.responses import Response, StreamingResponse from .auth import adjust_payment_for_tokens from .core import get_logger from .core.db import ApiKey, AsyncSession, create_session -from .core.settings import settings from .payment.helpers import create_error_response logger = get_logger(__name__) -def prepare_upstream_headers(request_headers: dict) -> dict: - upstream_api_key = settings.upstream_api_key - logger.debug( - "Preparing upstream headers", - extra={ - "original_headers_count": len(request_headers), - "has_upstream_api_key": bool(upstream_api_key), - }, - ) - - headers = dict(request_headers) - removed_headers = [] - - for header in [ - "host", - "content-length", - "refund-lnurl", - "key-expiry-time", - "x-cashu", - ]: - if headers.pop(header, None) is not None: - removed_headers.append(header) - - if upstream_api_key: - headers["Authorization"] = f"Bearer {upstream_api_key}" - if headers.pop("authorization", None) is not None: - removed_headers.append("authorization (replaced with upstream key)") - else: - for auth_header in ["Authorization", "authorization"]: - if headers.pop(auth_header, None) is not None: - removed_headers.append(auth_header) - - logger.debug( - "Headers prepared for upstream", - extra={ - "final_headers_count": len(headers), - "removed_headers": removed_headers, - "added_upstream_auth": bool(upstream_api_key), - }, - ) - - return headers - - -def prepare_upstream_params( - path: str, query_params: Mapping[str, str] | None -) -> dict[str, str]: - params: dict[str, str] = dict(query_params or {}) - chat_api_version = settings.chat_completions_api_version - if path.endswith("chat/completions") and chat_api_version: - params["api-version"] = chat_api_version - return params - - -def _extract_upstream_error_message(body_bytes: bytes) -> tuple[str, str | None]: - message: str = "Upstream request failed" - upstream_code: str | None = None - if not body_bytes: - return message, upstream_code - try: - data = json.loads(body_bytes) - if isinstance(data, dict): - err = data.get("error") - if isinstance(err, dict): - raw_msg = err.get("message") or err.get("detail") or err.get("error") - if isinstance(raw_msg, (str, int, float)): - message = str(raw_msg) - upstream_code_raw = err.get("code") or err.get("type") - if isinstance(upstream_code_raw, (str, int, float)): - upstream_code = str(upstream_code_raw) - elif "message" in data and isinstance(data["message"], (str, int, float)): - message = str(data["message"]) # type: ignore[arg-type] - elif "detail" in data and isinstance(data["detail"], (str, int, float)): - message = str(data["detail"]) # type: ignore[arg-type] - except Exception: - preview = body_bytes.decode("utf-8", errors="ignore").strip() - if preview: - message = preview[:500] - return message, upstream_code - - -async def map_upstream_error_response( - request: Request, - path: str, - upstream_response: httpx.Response, -) -> Response: - status_code = upstream_response.status_code - headers = dict(upstream_response.headers) - content_type = headers.get("content-type", "") - try: - body_bytes = await upstream_response.aread() - except Exception: - body_bytes = b"" - - message, upstream_code = _extract_upstream_error_message(body_bytes) - lowered_message = message.lower() - lowered_code = (upstream_code or "").lower() - - error_type = "upstream_error" - mapped_status = 502 - - if status_code in (400, 422): - error_type = "invalid_request_error" - mapped_status = 400 - elif status_code in (401, 403): - error_type = "upstream_auth_error" - mapped_status = 502 - elif status_code == 404: - if path.endswith("chat/completions"): - error_type = "invalid_model" - mapped_status = 400 - if not message or message == "Upstream request failed": - message = "Requested model is not available upstream" - elif "model" in lowered_message or "model" in lowered_code: - error_type = "invalid_model" - mapped_status = 400 - if not message or message == "Upstream request failed": - message = "Requested model is not available upstream" - else: - error_type = "upstream_error" - mapped_status = 502 - elif status_code == 429: - error_type = "rate_limit_exceeded" - mapped_status = 429 - elif status_code >= 500: - error_type = "upstream_error" - mapped_status = 502 - - logger.debug( - "Mapped upstream error", - extra={ - "path": path, - "upstream_status": status_code, - "mapped_status": mapped_status, - "error_type": error_type, - "upstream_content_type": content_type, - "message_preview": message[:200], - }, - ) - - return create_error_response(error_type, message, mapped_status, request=request) - - -async def handle_streaming_chat_completion( - response: httpx.Response, key: ApiKey, max_cost_for_model: int -) -> StreamingResponse: - logger.info( - "Processing streaming chat completion", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "key_balance": key.balance, - "response_status": response.status_code, - }, - ) - - async def stream_with_cost(max_cost_for_model: int) -> AsyncGenerator[bytes, None]: - stored_chunks: list[bytes] = [] - usage_finalized: bool = False - last_model_seen: str | None = None - - async def finalize_without_usage() -> bytes | None: - nonlocal usage_finalized - if usage_finalized: - return None - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - return None - try: - fallback: dict = { - "model": last_model_seen or "unknown", - "usage": None, - } - cost_data = await adjust_payment_for_tokens( - fresh_key, fallback, new_session, max_cost_for_model - ) - usage_finalized = True - logger.info( - "Finalized streaming payment without explicit usage", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "cost_data": cost_data, - "balance_after_adjustment": fresh_key.balance, - }, - ) - return f"data: {json.dumps({'cost': cost_data})}\n\n".encode() - except Exception as cost_error: - logger.error( - "Error finalizing payment without usage", - extra={ - "error": str(cost_error), - "error_type": type(cost_error).__name__, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - return None - - try: - async for chunk in response.aiter_bytes(): - stored_chunks.append(chunk) - try: - for part in re.split(b"data: ", chunk): - if not part or part.strip() in (b"[DONE]", b""): - continue - try: - obj = json.loads(part) - if isinstance(obj, dict) and obj.get("model"): - last_model_seen = str(obj.get("model")) - except json.JSONDecodeError: - pass - except Exception: - pass - - yield chunk - - logger.debug( - "Streaming completed, analyzing usage data", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "chunks_count": len(stored_chunks), - }, - ) - - for i in range(len(stored_chunks) - 1, -1, -1): - chunk = stored_chunks[i] - if not chunk: - continue - try: - events = re.split(b"data: ", chunk) - for event_data in events: - if not event_data or event_data.strip() in (b"[DONE]", b""): - continue - try: - data = json.loads(event_data) - if isinstance(data, dict) and data.get("model"): - last_model_seen = str(data.get("model")) - if isinstance(data, dict) and isinstance( - data.get("usage"), dict - ): - async with create_session() as new_session: - fresh_key = await new_session.get( - key.__class__, key.hashed_key - ) - if fresh_key: - try: - cost_data = await adjust_payment_for_tokens( - fresh_key, - data, - new_session, - max_cost_for_model, - ) - usage_finalized = True - logger.info( - "Token adjustment completed for streaming", - extra={ - "key_hash": key.hashed_key[:8] - + "...", - "cost_data": cost_data, - "balance_after_adjustment": fresh_key.balance, - }, - ) - yield f"data: {json.dumps({'cost': cost_data})}\n\n".encode() - except Exception as cost_error: - logger.error( - "Error adjusting payment for streaming tokens", - extra={ - "error": str(cost_error), - "error_type": type( - cost_error - ).__name__, - "key_hash": key.hashed_key[:8] - + "...", - }, - ) - break - except json.JSONDecodeError: - continue - except Exception as e: - logger.error( - "Error processing streaming response chunk", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - - if not usage_finalized: - maybe_cost_event = await finalize_without_usage() - if maybe_cost_event is not None: - yield maybe_cost_event - - except Exception as stream_error: - logger.warning( - "Streaming interrupted; finalizing without usage", - extra={ - "error": str(stream_error), - "error_type": type(stream_error).__name__, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - await finalize_without_usage() - raise - - return StreamingResponse( - stream_with_cost(max_cost_for_model), - status_code=response.status_code, - headers=dict(response.headers), - ) - - -async def handle_non_streaming_chat_completion( - response: httpx.Response, - key: ApiKey, - session: AsyncSession, - deducted_max_cost: int, -) -> Response: - logger.info( - "Processing non-streaming chat completion", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "key_balance": key.balance, - "response_status": response.status_code, - }, - ) - - try: - content = await response.aread() - response_json = json.loads(content) - - logger.debug( - "Parsed response JSON", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "model": response_json.get("model", "unknown"), - "has_usage": "usage" in response_json, - }, - ) - - cost_data = await adjust_payment_for_tokens( - key, response_json, session, deducted_max_cost - ) - response_json["cost"] = cost_data - - logger.info( - "Token adjustment completed for non-streaming", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "cost_data": cost_data, - "model": response_json.get("model", "unknown"), - "balance_after_adjustment": key.balance, - }, - ) - - allowed_headers = { - "content-type", - "cache-control", - "date", - "vary", - "access-control-allow-origin", - "access-control-allow-methods", - "access-control-allow-headers", - "access-control-allow-credentials", - "access-control-expose-headers", - "access-control-max-age", - } - - response_headers = { - k: v for k, v in response.headers.items() if k.lower() in allowed_headers - } - - 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: - logger.error( - "Failed to parse JSON from upstream response", - extra={ - "error": str(e), - "key_hash": key.hashed_key[:8] + "...", - "content_preview": content[:200].decode(errors="ignore") - if content - else "empty", - }, - ) - raise - except Exception as e: - logger.error( - "Error processing non-streaming chat completion", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - raise - - class UpstreamProvider: def __init__( self, base_url: str, api_key: str, chat_completions_api_version: str = "" @@ -482,6 +80,362 @@ class UpstreamProvider: params["api-version"] = self.chat_completions_api_version return params + def _extract_upstream_error_message( + self, body_bytes: bytes + ) -> tuple[str, str | None]: + message: str = "Upstream request failed" + upstream_code: str | None = None + if not body_bytes: + return message, upstream_code + try: + data = json.loads(body_bytes) + if isinstance(data, dict): + err = data.get("error") + if isinstance(err, dict): + raw_msg = ( + err.get("message") or err.get("detail") or err.get("error") + ) + if isinstance(raw_msg, (str, int, float)): + message = str(raw_msg) + upstream_code_raw = err.get("code") or err.get("type") + if isinstance(upstream_code_raw, (str, int, float)): + upstream_code = str(upstream_code_raw) + elif "message" in data and isinstance( + data["message"], (str, int, float) + ): + message = str(data["message"]) # type: ignore[arg-type] + elif "detail" in data and isinstance(data["detail"], (str, int, float)): + message = str(data["detail"]) # type: ignore[arg-type] + except Exception: + preview = body_bytes.decode("utf-8", errors="ignore").strip() + if preview: + message = preview[:500] + return message, upstream_code + + async def map_upstream_error_response( + self, request: Request, path: str, upstream_response: httpx.Response + ) -> Response: + status_code = upstream_response.status_code + headers = dict(upstream_response.headers) + content_type = headers.get("content-type", "") + try: + body_bytes = await upstream_response.aread() + except Exception: + body_bytes = b"" + + message, upstream_code = self._extract_upstream_error_message(body_bytes) + lowered_message = message.lower() + lowered_code = (upstream_code or "").lower() + + error_type = "upstream_error" + mapped_status = 502 + + if status_code in (400, 422): + error_type = "invalid_request_error" + mapped_status = 400 + elif status_code in (401, 403): + error_type = "upstream_auth_error" + mapped_status = 502 + elif status_code == 404: + if path.endswith("chat/completions"): + error_type = "invalid_model" + mapped_status = 400 + if not message or message == "Upstream request failed": + message = "Requested model is not available upstream" + elif "model" in lowered_message or "model" in lowered_code: + error_type = "invalid_model" + mapped_status = 400 + if not message or message == "Upstream request failed": + message = "Requested model is not available upstream" + else: + error_type = "upstream_error" + mapped_status = 502 + elif status_code == 429: + error_type = "rate_limit_exceeded" + mapped_status = 429 + elif status_code >= 500: + error_type = "upstream_error" + mapped_status = 502 + + logger.debug( + "Mapped upstream error", + extra={ + "path": path, + "upstream_status": status_code, + "mapped_status": mapped_status, + "error_type": error_type, + "upstream_content_type": content_type, + "message_preview": message[:200], + }, + ) + + return create_error_response( + error_type, message, mapped_status, request=request + ) + + async def handle_streaming_chat_completion( + self, response: httpx.Response, key: ApiKey, max_cost_for_model: int + ) -> StreamingResponse: + logger.info( + "Processing streaming chat completion", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "key_balance": key.balance, + "response_status": response.status_code, + }, + ) + + async def stream_with_cost( + max_cost_for_model: int, + ) -> AsyncGenerator[bytes, None]: + stored_chunks: list[bytes] = [] + usage_finalized: bool = False + last_model_seen: str | None = None + + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + return None + try: + fallback: dict = { + "model": last_model_seen or "unknown", + "usage": None, + } + cost_data = await adjust_payment_for_tokens( + fresh_key, fallback, new_session, max_cost_for_model + ) + usage_finalized = True + logger.info( + "Finalized streaming payment without explicit usage", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "cost_data": cost_data, + "balance_after_adjustment": fresh_key.balance, + }, + ) + return f"data: {json.dumps({'cost': cost_data})}\n\n".encode() + except Exception as cost_error: + logger.error( + "Error finalizing payment without usage", + extra={ + "error": str(cost_error), + "error_type": type(cost_error).__name__, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + return None + + try: + async for chunk in response.aiter_bytes(): + stored_chunks.append(chunk) + try: + for part in re.split(b"data: ", chunk): + if not part or part.strip() in (b"[DONE]", b""): + continue + try: + obj = json.loads(part) + if isinstance(obj, dict) and obj.get("model"): + last_model_seen = str(obj.get("model")) + except json.JSONDecodeError: + pass + except Exception: + pass + + yield chunk + + logger.debug( + "Streaming completed, analyzing usage data", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "chunks_count": len(stored_chunks), + }, + ) + + for i in range(len(stored_chunks) - 1, -1, -1): + chunk = stored_chunks[i] + if not chunk: + continue + try: + events = re.split(b"data: ", chunk) + for event_data in events: + if not event_data or event_data.strip() in (b"[DONE]", b""): + continue + try: + data = json.loads(event_data) + if isinstance(data, dict) and data.get("model"): + last_model_seen = str(data.get("model")) + if isinstance(data, dict) and isinstance( + data.get("usage"), dict + ): + async with create_session() as new_session: + fresh_key = await new_session.get( + key.__class__, key.hashed_key + ) + if fresh_key: + try: + cost_data = ( + await adjust_payment_for_tokens( + fresh_key, + data, + new_session, + max_cost_for_model, + ) + ) + usage_finalized = True + logger.info( + "Token adjustment completed for streaming", + extra={ + "key_hash": key.hashed_key[:8] + + "...", + "cost_data": cost_data, + "balance_after_adjustment": fresh_key.balance, + }, + ) + yield f"data: {json.dumps({'cost': cost_data})}\n\n".encode() + except Exception as cost_error: + logger.error( + "Error adjusting payment for streaming tokens", + extra={ + "error": str(cost_error), + "error_type": type( + cost_error + ).__name__, + "key_hash": key.hashed_key[:8] + + "...", + }, + ) + break + except json.JSONDecodeError: + continue + except Exception as e: + logger.error( + "Error processing streaming response chunk", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + if not usage_finalized: + maybe_cost_event = await finalize_without_usage() + if maybe_cost_event is not None: + yield maybe_cost_event + + except Exception as stream_error: + logger.warning( + "Streaming interrupted; finalizing without usage", + extra={ + "error": str(stream_error), + "error_type": type(stream_error).__name__, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + await finalize_without_usage() + raise + + return StreamingResponse( + stream_with_cost(max_cost_for_model), + status_code=response.status_code, + headers=dict(response.headers), + ) + + async def handle_non_streaming_chat_completion( + self, + response: httpx.Response, + key: ApiKey, + session: AsyncSession, + deducted_max_cost: int, + ) -> Response: + logger.info( + "Processing non-streaming chat completion", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "key_balance": key.balance, + "response_status": response.status_code, + }, + ) + + try: + content = await response.aread() + response_json = json.loads(content) + + logger.debug( + "Parsed response JSON", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "model": response_json.get("model", "unknown"), + "has_usage": "usage" in response_json, + }, + ) + + cost_data = await adjust_payment_for_tokens( + key, response_json, session, deducted_max_cost + ) + response_json["cost"] = cost_data + + logger.info( + "Token adjustment completed for non-streaming", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "cost_data": cost_data, + "model": response_json.get("model", "unknown"), + "balance_after_adjustment": key.balance, + }, + ) + + allowed_headers = { + "content-type", + "cache-control", + "date", + "vary", + "access-control-allow-origin", + "access-control-allow-methods", + "access-control-allow-headers", + "access-control-allow-credentials", + "access-control-expose-headers", + "access-control-max-age", + } + + response_headers = { + k: v + for k, v in response.headers.items() + if k.lower() in allowed_headers + } + + 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: + logger.error( + "Failed to parse JSON from upstream response", + extra={ + "error": str(e), + "key_hash": key.hashed_key[:8] + "...", + "content_preview": content[:200].decode(errors="ignore") + if content + else "empty", + }, + ) + raise + except Exception as e: + logger.error( + "Error processing non-streaming chat completion", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + raise + async def forward_request( self, request: Request, @@ -491,15 +445,6 @@ class UpstreamProvider: key: ApiKey, max_cost_for_model: int, session: AsyncSession, - map_upstream_error_response: Callable[ - [Request, str, httpx.Response], Awaitable[Response] - ], - handle_streaming_chat_completion: Callable[ - [httpx.Response, ApiKey, int], Awaitable[StreamingResponse] - ], - handle_non_streaming_chat_completion: Callable[ - [httpx.Response, ApiKey, AsyncSession, int], Awaitable[Response] - ], ) -> Response | StreamingResponse: if path.startswith("v1/"): path = path.replace("v1/", "") @@ -559,7 +504,7 @@ class UpstreamProvider: if response.status_code != 200: try: - mapped_error = await map_upstream_error_response( + mapped_error = await self.map_upstream_error_response( request, path, response ) finally: @@ -602,7 +547,7 @@ class UpstreamProvider: ) if is_streaming and response.status_code == 200: - result = await handle_streaming_chat_completion( + result = await self.handle_streaming_chat_completion( response, key, max_cost_for_model ) background_tasks = BackgroundTasks() @@ -613,7 +558,7 @@ class UpstreamProvider: elif response.status_code == 200: try: - return await handle_non_streaming_chat_completion( + return await self.handle_non_streaming_chat_completion( response, key, session, max_cost_for_model ) finally: @@ -701,9 +646,6 @@ class UpstreamProvider: request: Request, path: str, headers: dict, - map_upstream_error_response: Callable[ - [Request, str, httpx.Response], Awaitable[Response] - ], ) -> Response | StreamingResponse: if path.startswith("v1/"): path = path.replace("v1/", "") @@ -736,7 +678,7 @@ class UpstreamProvider: ) if response.status_code != 200: try: - mapped = await map_upstream_error_response( + mapped = await self.map_upstream_error_response( request, path, response ) finally: @@ -769,6 +711,421 @@ class UpstreamProvider: request=request, ) + async def get_x_cashu_cost( + self, response_data: dict, max_cost_for_model: int + ) -> MaxCostData | CostData | None: + from .payment.cost_caculation import ( + CostData, + CostDataError, + MaxCostData, + calculate_cost, + ) + + model = response_data.get("model", None) + logger.debug( + "Calculating cost for response", + extra={"model": model, "has_usage": "usage" in response_data}, + ) + + async with create_session() as session: + match await calculate_cost(response_data, max_cost_for_model, session): + case MaxCostData() as cost: + logger.debug( + "Using max cost pricing", + extra={"model": model, "max_cost_msats": cost.total_msats}, + ) + return cost + case CostData() as cost: + logger.debug( + "Using token-based pricing", + extra={ + "model": model, + "total_cost_msats": cost.total_msats, + "input_msats": cost.input_msats, + "output_msats": cost.output_msats, + }, + ) + return cost + case CostDataError() as error: + logger.error( + "Cost calculation error", + extra={ + "model": model, + "error_message": error.message, + "error_code": error.code, + }, + ) + raise HTTPException( + status_code=400, + detail={ + "error": { + "message": error.message, + "type": "invalid_request_error", + "code": error.code, + } + }, + ) + return None + + async def send_refund(self, amount: int, unit: str, mint: str | None = None) -> str: + from .wallet import send_token + + logger.debug( + "Creating refund token", + extra={"amount": amount, "unit": unit, "mint": mint}, + ) + + max_retries = 3 + last_exception = None + + for attempt in range(max_retries): + try: + refund_token = await send_token(amount, unit=unit, mint_url=mint) + + logger.info( + "Refund token created successfully", + extra={ + "amount": amount, + "unit": unit, + "mint": mint, + "attempt": attempt + 1, + "token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, + }, + ) + + return refund_token + except Exception as e: + last_exception = e + if attempt < max_retries - 1: + logger.warning( + "Refund token creation failed, retrying", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "attempt": attempt + 1, + "max_retries": max_retries, + "amount": amount, + "unit": unit, + "mint": mint, + }, + ) + else: + logger.error( + "Failed to create refund token after all retries", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "attempt": attempt + 1, + "max_retries": max_retries, + "amount": amount, + "unit": unit, + "mint": mint, + }, + ) + + raise HTTPException( + status_code=401, + detail={ + "error": { + "message": f"failed to create refund after {max_retries} attempts: {str(last_exception)}", + "type": "invalid_request_error", + "code": "send_token_failed", + } + }, + ) + + async def handle_x_cashu_streaming_response( + self, + content_str: str, + response: httpx.Response, + amount: int, + unit: str, + max_cost_for_model: int, + ) -> StreamingResponse: + logger.debug( + "Processing streaming response", + extra={ + "amount": amount, + "unit": unit, + "content_lines": len(content_str.strip().split("\n")), + }, + ) + + response_headers = dict(response.headers) + if "transfer-encoding" in response_headers: + del response_headers["transfer-encoding"] + if "content-encoding" in response_headers: + del response_headers["content-encoding"] + + usage_data = None + model = None + + lines = content_str.strip().split("\n") + for line in lines: + if line.startswith("data: "): + try: + data_json = json.loads(line[6:]) + if "usage" in data_json: + usage_data = data_json["usage"] + model = data_json.get("model") + elif "model" in data_json and not model: + model = data_json["model"] + except json.JSONDecodeError: + continue + + if usage_data and model: + logger.debug( + "Found usage data in streaming response", + extra={ + "model": model, + "usage_data": usage_data, + "amount": amount, + "unit": unit, + }, + ) + + response_data = {"usage": usage_data, "model": model} + try: + cost_data = await self.get_x_cashu_cost( + response_data, max_cost_for_model + ) + if cost_data: + if unit == "msat": + refund_amount = amount - cost_data.total_msats + elif unit == "sat": + refund_amount = amount - (cost_data.total_msats + 999) // 1000 + else: + raise ValueError(f"Invalid unit: {unit}") + + if refund_amount > 0: + logger.info( + "Processing refund for streaming response", + extra={ + "original_amount": amount, + "cost_msats": cost_data.total_msats, + "refund_amount": refund_amount, + "unit": unit, + "model": model, + }, + ) + + refund_token = await self.send_refund(refund_amount, unit) + response_headers["X-Cashu"] = refund_token + + logger.info( + "Refund processed for streaming response", + extra={ + "refund_amount": refund_amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, + }, + ) + else: + logger.debug( + "No refund needed for streaming response", + extra={ + "amount": amount, + "cost_msats": cost_data.total_msats, + "model": model, + }, + ) + except Exception as e: + logger.error( + "Error calculating cost for streaming response", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "model": model, + "amount": amount, + "unit": unit, + }, + ) + + async def generate() -> AsyncGenerator[bytes, None]: + for line in lines: + yield (line + "\n").encode("utf-8") + + return StreamingResponse( + generate(), + status_code=response.status_code, + headers=response_headers, + media_type="text/plain", + ) + + async def handle_x_cashu_non_streaming_response( + self, + content_str: str, + response: httpx.Response, + amount: int, + unit: str, + max_cost_for_model: int, + ) -> Response: + logger.debug( + "Processing non-streaming response", + extra={"amount": amount, "unit": unit, "content_length": len(content_str)}, + ) + + try: + response_json = json.loads(content_str) + cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model) + + if not cost_data: + logger.error( + "Failed to calculate cost for response", + extra={ + "amount": amount, + "unit": unit, + "response_model": response_json.get("model", "unknown"), + }, + ) + return Response( + content=json.dumps( + { + "error": { + "message": "Error forwarding request to upstream", + "type": "upstream_error", + "code": response.status_code, + } + } + ), + status_code=response.status_code, + media_type="application/json", + ) + + response_headers = dict(response.headers) + if "transfer-encoding" in response_headers: + del response_headers["transfer-encoding"] + if "content-encoding" in response_headers: + del response_headers["content-encoding"] + + if unit == "msat": + refund_amount = amount - cost_data.total_msats + elif unit == "sat": + refund_amount = amount - (cost_data.total_msats + 999) // 1000 + else: + raise ValueError(f"Invalid unit: {unit}") + + logger.info( + "Processing non-streaming response cost calculation", + extra={ + "original_amount": amount, + "cost_msats": cost_data.total_msats, + "refund_amount": refund_amount, + "unit": unit, + "model": response_json.get("model", "unknown"), + }, + ) + + if refund_amount > 0: + refund_token = await self.send_refund(refund_amount, unit) + response_headers["X-Cashu"] = refund_token + + logger.info( + "Refund processed for non-streaming response", + extra={ + "refund_amount": refund_amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, + }, + ) + + return Response( + content=content_str, + status_code=response.status_code, + headers=response_headers, + media_type="application/json", + ) + except json.JSONDecodeError as e: + logger.error( + "Failed to parse JSON from upstream response", + extra={ + "error": str(e), + "content_preview": content_str[:200] + "..." + if len(content_str) > 200 + else content_str, + "amount": amount, + "unit": unit, + }, + ) + + from .wallet import send_token + + emergency_refund = amount + refund_token = await send_token(emergency_refund, unit=unit) + response.headers["X-Cashu"] = refund_token + + logger.warning( + "Emergency refund issued due to JSON parse error", + extra={ + "original_amount": amount, + "refund_amount": emergency_refund, + "deduction": 60, + }, + ) + + return Response( + content=content_str, + status_code=response.status_code, + headers=dict(response.headers), + media_type="application/json", + ) + + async def handle_x_cashu_chat_completion( + self, response: httpx.Response, amount: int, unit: str, max_cost_for_model: int + ) -> StreamingResponse | Response: + logger.debug( + "Handling chat completion response", + extra={"amount": amount, "unit": unit, "status_code": response.status_code}, + ) + + try: + content = await response.aread() + content_str = ( + content.decode("utf-8") if isinstance(content, bytes) else content + ) + is_streaming = content_str.startswith("data:") or "data:" in content_str + + logger.debug( + "Chat completion response analysis", + extra={ + "is_streaming": is_streaming, + "content_length": len(content_str), + "amount": amount, + "unit": unit, + }, + ) + + if is_streaming: + return await self.handle_x_cashu_streaming_response( + content_str, response, amount, unit, max_cost_for_model + ) + else: + return await self.handle_x_cashu_non_streaming_response( + content_str, response, amount, unit, max_cost_for_model + ) + + except Exception as e: + logger.error( + "Error processing chat completion response", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "amount": amount, + "unit": unit, + }, + ) + return StreamingResponse( + response.aiter_bytes(), + status_code=response.status_code, + headers=dict(response.headers), + ) + async def forward_x_cashu_request( self, request: Request, @@ -777,10 +1134,6 @@ class UpstreamProvider: amount: int, unit: str, max_cost_for_model: int, - handle_x_cashu_chat_completion: Callable[ - [httpx.Response, int, str, int], Awaitable[Response | StreamingResponse] - ], - send_refund: Callable[[int, str], Awaitable[str]], ) -> Response | StreamingResponse: if path.startswith("v1/"): path = path.replace("v1/", "") @@ -834,7 +1187,7 @@ class UpstreamProvider: }, ) - refund_token = await send_refund(amount - 60, unit) + refund_token = await self.send_refund(amount - 60, unit) logger.info( "Refund processed for failed upstream request", @@ -871,7 +1224,7 @@ class UpstreamProvider: extra={"path": path, "amount": amount, "unit": unit}, ) - result = await handle_x_cashu_chat_completion( + result = await self.handle_x_cashu_chat_completion( response, amount, unit, max_cost_for_model ) background_tasks = BackgroundTasks() @@ -948,8 +1301,6 @@ class UpstreamProvider: amount, unit, max_cost_for_model, - handle_x_cashu_chat_completion, - send_refund, ) except Exception as e: error_message = str(e) @@ -997,416 +1348,3 @@ class UpstreamProvider: request=request, token=x_cashu_token, ) - - -async def handle_x_cashu_chat_completion( - response: httpx.Response, amount: int, unit: str, max_cost_for_model: int -) -> StreamingResponse | Response: - logger.debug( - "Handling chat completion response", - extra={"amount": amount, "unit": unit, "status_code": response.status_code}, - ) - - try: - content = await response.aread() - content_str = content.decode("utf-8") if isinstance(content, bytes) else content - is_streaming = content_str.startswith("data:") or "data:" in content_str - - logger.debug( - "Chat completion response analysis", - extra={ - "is_streaming": is_streaming, - "content_length": len(content_str), - "amount": amount, - "unit": unit, - }, - ) - - if is_streaming: - return await handle_x_cashu_streaming_response( - content_str, response, amount, unit, max_cost_for_model - ) - else: - return await handle_x_cashu_non_streaming_response( - content_str, response, amount, unit, max_cost_for_model - ) - - except Exception as e: - logger.error( - "Error processing chat completion response", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "amount": amount, - "unit": unit, - }, - ) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - ) - - -async def handle_x_cashu_streaming_response( - content_str: str, - response: httpx.Response, - amount: int, - unit: str, - max_cost_for_model: int, -) -> StreamingResponse: - logger.debug( - "Processing streaming response", - extra={ - "amount": amount, - "unit": unit, - "content_lines": len(content_str.strip().split("\n")), - }, - ) - - response_headers = dict(response.headers) - if "transfer-encoding" in response_headers: - del response_headers["transfer-encoding"] - if "content-encoding" in response_headers: - del response_headers["content-encoding"] - - usage_data = None - model = None - - lines = content_str.strip().split("\n") - for line in lines: - if line.startswith("data: "): - try: - data_json = json.loads(line[6:]) - if "usage" in data_json: - usage_data = data_json["usage"] - model = data_json.get("model") - elif "model" in data_json and not model: - model = data_json["model"] - except json.JSONDecodeError: - continue - - if usage_data and model: - logger.debug( - "Found usage data in streaming response", - extra={ - "model": model, - "usage_data": usage_data, - "amount": amount, - "unit": unit, - }, - ) - - response_data = {"usage": usage_data, "model": model} - try: - cost_data = await get_x_cashu_cost(response_data, max_cost_for_model) - if cost_data: - if unit == "msat": - refund_amount = amount - cost_data.total_msats - elif unit == "sat": - refund_amount = amount - (cost_data.total_msats + 999) // 1000 - else: - raise ValueError(f"Invalid unit: {unit}") - - if refund_amount > 0: - logger.info( - "Processing refund for streaming response", - extra={ - "original_amount": amount, - "cost_msats": cost_data.total_msats, - "refund_amount": refund_amount, - "unit": unit, - "model": model, - }, - ) - - refund_token = await send_refund(refund_amount, unit) - response_headers["X-Cashu"] = refund_token - - logger.info( - "Refund processed for streaming response", - extra={ - "refund_amount": refund_amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - else: - logger.debug( - "No refund needed for streaming response", - extra={ - "amount": amount, - "cost_msats": cost_data.total_msats, - "model": model, - }, - ) - except Exception as e: - logger.error( - "Error calculating cost for streaming response", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "model": model, - "amount": amount, - "unit": unit, - }, - ) - - async def generate() -> AsyncGenerator[bytes, None]: - for line in lines: - yield (line + "\n").encode("utf-8") - - return StreamingResponse( - generate(), - status_code=response.status_code, - headers=response_headers, - media_type="text/plain", - ) - - -async def handle_x_cashu_non_streaming_response( - content_str: str, - response: httpx.Response, - amount: int, - unit: str, - max_cost_for_model: int, -) -> Response: - logger.debug( - "Processing non-streaming response", - extra={"amount": amount, "unit": unit, "content_length": len(content_str)}, - ) - - try: - response_json = json.loads(content_str) - cost_data = await get_x_cashu_cost(response_json, max_cost_for_model) - - if not cost_data: - logger.error( - "Failed to calculate cost for response", - extra={ - "amount": amount, - "unit": unit, - "response_model": response_json.get("model", "unknown"), - }, - ) - return Response( - content=json.dumps( - { - "error": { - "message": "Error forwarding request to upstream", - "type": "upstream_error", - "code": response.status_code, - } - } - ), - status_code=response.status_code, - media_type="application/json", - ) - - response_headers = dict(response.headers) - if "transfer-encoding" in response_headers: - del response_headers["transfer-encoding"] - if "content-encoding" in response_headers: - del response_headers["content-encoding"] - - if unit == "msat": - refund_amount = amount - cost_data.total_msats - elif unit == "sat": - refund_amount = amount - (cost_data.total_msats + 999) // 1000 - else: - raise ValueError(f"Invalid unit: {unit}") - - logger.info( - "Processing non-streaming response cost calculation", - extra={ - "original_amount": amount, - "cost_msats": cost_data.total_msats, - "refund_amount": refund_amount, - "unit": unit, - "model": response_json.get("model", "unknown"), - }, - ) - - if refund_amount > 0: - refund_token = await send_refund(refund_amount, unit) - response_headers["X-Cashu"] = refund_token - - logger.info( - "Refund processed for non-streaming response", - extra={ - "refund_amount": refund_amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - - return Response( - content=content_str, - status_code=response.status_code, - headers=response_headers, - media_type="application/json", - ) - except json.JSONDecodeError as e: - logger.error( - "Failed to parse JSON from upstream response", - extra={ - "error": str(e), - "content_preview": content_str[:200] + "..." - if len(content_str) > 200 - else content_str, - "amount": amount, - "unit": unit, - }, - ) - - from .wallet import send_token - - emergency_refund = amount - refund_token = await send_token(emergency_refund, unit=unit) - response.headers["X-Cashu"] = refund_token - - logger.warning( - "Emergency refund issued due to JSON parse error", - extra={ - "original_amount": amount, - "refund_amount": emergency_refund, - "deduction": 60, - }, - ) - - return Response( - content=content_str, - status_code=response.status_code, - headers=dict(response.headers), - media_type="application/json", - ) - - -async def get_x_cashu_cost( - response_data: dict, max_cost_for_model: int -) -> MaxCostData | CostData | None: - from .payment.cost_caculation import ( - CostData, - CostDataError, - MaxCostData, - calculate_cost, - ) - - model = response_data.get("model", None) - logger.debug( - "Calculating cost for response", - extra={"model": model, "has_usage": "usage" in response_data}, - ) - - async with create_session() as session: - match await calculate_cost(response_data, max_cost_for_model, session): - case MaxCostData() as cost: - logger.debug( - "Using max cost pricing", - extra={"model": model, "max_cost_msats": cost.total_msats}, - ) - return cost - case CostData() as cost: - logger.debug( - "Using token-based pricing", - extra={ - "model": model, - "total_cost_msats": cost.total_msats, - "input_msats": cost.input_msats, - "output_msats": cost.output_msats, - }, - ) - return cost - case CostDataError() as error: - logger.error( - "Cost calculation error", - extra={ - "model": model, - "error_message": error.message, - "error_code": error.code, - }, - ) - raise HTTPException( - status_code=400, - detail={ - "error": { - "message": error.message, - "type": "invalid_request_error", - "code": error.code, - } - }, - ) - return None - - -async def send_refund(amount: int, unit: str, mint: str | None = None) -> str: - from .wallet import send_token - - logger.debug( - "Creating refund token", extra={"amount": amount, "unit": unit, "mint": mint} - ) - - max_retries = 3 - last_exception = None - - for attempt in range(max_retries): - try: - refund_token = await send_token(amount, unit=unit, mint_url=mint) - - logger.info( - "Refund token created successfully", - extra={ - "amount": amount, - "unit": unit, - "mint": mint, - "attempt": attempt + 1, - "token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - - return refund_token - except Exception as e: - last_exception = e - if attempt < max_retries - 1: - logger.warning( - "Refund token creation failed, retrying", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "attempt": attempt + 1, - "max_retries": max_retries, - "amount": amount, - "unit": unit, - "mint": mint, - }, - ) - else: - logger.error( - "Failed to create refund token after all retries", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "attempt": attempt + 1, - "max_retries": max_retries, - "amount": amount, - "unit": unit, - "mint": mint, - }, - ) - - raise HTTPException( - status_code=401, - detail={ - "error": { - "message": f"failed to create refund after {max_retries} attempts: {str(last_exception)}", - "type": "invalid_request_error", - "code": "send_token_failed", - } - }, - )