From f65785924938947c962ef09f73b78c98615ebca9 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 4 Oct 2025 17:38:00 +0800 Subject: [PATCH 01/95] refactor upstream functions into UpstreamProvider class --- routstr/payment/helpers.py | 59 -- routstr/payment/x_cashu.py | 664 ----------------- routstr/proxy.py | 698 +----------------- routstr/upstream.py | 1412 ++++++++++++++++++++++++++++++++++++ 4 files changed, 1443 insertions(+), 1390 deletions(-) delete mode 100644 routstr/payment/x_cashu.py create mode 100644 routstr/upstream.py diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 6dc4b8ff..37c7c8be 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -1,6 +1,5 @@ import json import math -from typing import Mapping from fastapi import HTTPException, Response from fastapi.requests import Request @@ -257,61 +256,3 @@ def create_error_response( media_type="application/json", headers={"X-Cashu": token} if token else {}, ) - - -def prepare_upstream_headers(request_headers: dict) -> dict: - """Prepare headers for upstream request, removing sensitive/problematic ones.""" - 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) - - # Remove headers that shouldn't be forwarded - 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) - - # Handle authorization - 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]: - """Prepare query params for upstream request, optionally adding api-version for chat/completions.""" - 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 diff --git a/routstr/payment/x_cashu.py b/routstr/payment/x_cashu.py deleted file mode 100644 index f671fbd8..00000000 --- a/routstr/payment/x_cashu.py +++ /dev/null @@ -1,664 +0,0 @@ -import json -import traceback -from typing import AsyncGenerator - -import httpx -from fastapi import BackgroundTasks, HTTPException, Request -from fastapi.responses import Response, StreamingResponse - -from ..core import get_logger -from ..core.db import create_session -from ..core.settings import settings -from ..wallet import recieve_token, send_token -from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost -from .helpers import ( - create_error_response, - prepare_upstream_headers, - prepare_upstream_params, -) - -logger = get_logger(__name__) - - -async def x_cashu_handler( - request: Request, x_cashu_token: str, path: str, max_cost_for_model: int -) -> Response | StreamingResponse: - """Handle X-Cashu token payment requests.""" - logger.info( - "Processing X-Cashu payment request", - extra={ - "path": path, - "method": request.method, - "token_preview": x_cashu_token[:20] + "..." - if len(x_cashu_token) > 20 - else x_cashu_token, - }, - ) - - try: - headers = dict(request.headers) - amount, unit, mint = await recieve_token(x_cashu_token) - headers = prepare_upstream_headers(dict(request.headers)) - - logger.info( - "X-Cashu token redeemed successfully", - extra={"amount": amount, "unit": unit, "path": path, "mint": mint}, - ) - - return await forward_to_upstream( - request, path, headers, amount, unit, max_cost_for_model - ) - except Exception as e: - error_message = str(e) - logger.error( - "X-Cashu payment request failed", - extra={ - "error": error_message, - "error_type": type(e).__name__, - "path": path, - "method": request.method, - }, - ) - - # Handle specific CASHU errors with appropriate HTTP status codes - if "already spent" in error_message.lower(): - return create_error_response( - "token_already_spent", - "The provided CASHU token has already been spent", - 400, - request=request, - token=x_cashu_token, - ) - - if "invalid token" in error_message.lower(): - return create_error_response( - "invalid_token", - "The provided CASHU token is invalid", - 400, - request=request, - token=x_cashu_token, - ) - - if "mint error" in error_message.lower(): - return create_error_response( - "mint_error", - f"CASHU mint error: {error_message}", - 422, - request=request, - token=x_cashu_token, - ) - - # Generic error for other cases - return create_error_response( - "cashu_error", - f"CASHU token processing failed: {error_message}", - 400, - request=request, - token=x_cashu_token, - ) - - -async def forward_to_upstream( - 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/"): - path = path.replace("v1/", "") - - url = f"{settings.upstream_base_url}/{path}" - - logger.debug( - "Forwarding request to upstream", - extra={ - "url": url, - "method": request.method, - "path": path, - "amount": amount, - "unit": unit, - }, - ) - - async with httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) as client: - try: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=request.stream(), - params=prepare_upstream_params(path, request.query_params), - ), - stream=True, - ) - - logger.debug( - "Received upstream response", - extra={ - "status_code": response.status_code, - "path": path, - "response_headers": dict(response.headers), - }, - ) - - if response.status_code != 200: - logger.warning( - "Upstream request failed, processing refund", - extra={ - "status_code": response.status_code, - "path": path, - "amount": amount, - "unit": unit, - }, - ) - - refund_token = await send_refund(amount - 60, unit) - - logger.info( - "Refund processed for failed upstream request", - extra={ - "status_code": response.status_code, - "refund_amount": amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - - error_response = Response( - content=json.dumps( - { - "error": { - "message": "Error forwarding request to upstream", - "type": "upstream_error", - "code": response.status_code, - "refund_token": refund_token, - } - } - ), - status_code=response.status_code, - media_type="application/json", - ) - error_response.headers["X-Cashu"] = refund_token - return error_response - - if path.endswith("chat/completions"): - logger.debug( - "Processing chat completion response", - extra={"path": path, "amount": amount, "unit": 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 - return result - - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - - logger.debug( - "Streaming non-chat response", - extra={"path": path, "status_code": response.status_code}, - ) - - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, - ) - except Exception as exc: - tb = traceback.format_exc() - logger.error( - "Unexpected error in upstream forwarding", - extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, - "path": path, - "query_params": dict(request.query_params), - "traceback": tb, - }, - ) - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, - ) - - -async def handle_x_cashu_chat_completion( - 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( - "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_streaming_response( - content_str, response, amount, unit, max_cost_for_model - ) - else: - return await handle_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 the original response if we can't process it - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - ) - - -async def handle_streaming_response( - 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( - "Processing streaming response", - extra={ - "amount": amount, - "unit": unit, - "content_lines": len(content_str.strip().split("\n")), - }, - ) - - # Initialize response headers early so they can be modified during processing - 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"] - - # For streaming responses, we'll extract the final usage data - # and calculate cost based on that - usage_data = None - model = None - - # Parse SSE format to extract usage information - lines = content_str.strip().split("\n") - for line in lines: - if line.startswith("data: "): - try: - data_json = json.loads(line[6:]) # Remove 'data: ' prefix - # Look for usage information in the final chunks - 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 - - response_headers = dict(response.headers) - # If we found usage data, calculate cost and refund - 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_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_non_streaming_response( - content_str: str, - response: httpx.Response, - amount: int, - unit: str, - max_cost_for_model: int, -) -> Response: - """Handle regular JSON 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_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, - }, - ) - - # Emergency refund with small deduction for processing - 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 original content if JSON parsing fails - return Response( - content=content_str, - status_code=response.status_code, - headers=dict(response.headers), - media_type="application/json", - ) - - -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", 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: - """Send a refund using Cashu tokens.""" - 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, - }, - ) - - # If we get here, all retries failed - 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", - } - }, - ) diff --git a/routstr/proxy.py b/routstr/proxy.py index 89660837..abee68b7 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -1,611 +1,33 @@ import json -import re -import traceback -from typing import AsyncGenerator -import httpx -from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request +from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import Response, StreamingResponse -from .auth import ( - adjust_payment_for_tokens, - pay_for_request, - revert_pay_for_request, - validate_bearer_key, -) +from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key from .core import get_logger -from .core.db import ApiKey, AsyncSession, create_session, get_session +from .core.db import ApiKey, AsyncSession, get_session from .core.settings import settings from .payment.helpers import ( calculate_discounted_max_cost, check_token_balance, create_error_response, get_max_cost_for_model, - prepare_upstream_headers, - prepare_upstream_params, ) -from .payment.x_cashu import x_cashu_handler +from .upstream import ( + UpstreamProvider, + handle_non_streaming_chat_completion, + handle_streaming_chat_completion, + map_upstream_error_response, +) logger = get_logger(__name__) proxy_router = APIRouter() - -def _extract_upstream_error_message(body_bytes: bytes) -> tuple[str, str | None]: - """Extract a human-friendly message and optional upstream error code from a response body.""" - 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: - """Map upstream non-200 responses to standardized error responses. - - - Known cases are mapped to friendly messages and appropriate status codes - - Unknown errors are converted to a generic 502 - """ - 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 - - # Specific mappings - 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: - # Many providers return 404 for unknown models or routes - 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 - - # Include upstream content type hint in logs for diagnostics - 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: - """Handle streaming chat completion responses with token-based pricing.""" - 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) - # Opportunistically capture model id - 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), - }, - ) - - # Process stored chunks to find usage data from the tail - 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 we reach here without finding usage, finalize with max-cost - 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: - # On stream interruption, still finalize reservation with max-cost - 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: - """Handle non-streaming chat completion responses with token-based pricing.""" - 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, - }, - ) - - # Keep only standard headers that are safe to pass through - 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_to_upstream( - request: Request, - path: str, - headers: dict, - request_body: bytes | None, - key: ApiKey, - max_cost_for_model: int, - session: AsyncSession, -) -> Response | StreamingResponse: - """Forward request to upstream and handle the response.""" - if path.startswith("v1/"): - path = path.replace("v1/", "") - - url = f"{settings.upstream_base_url}/{path}" - - logger.info( - "Forwarding request to upstream", - extra={ - "url": url, - "method": request.method, - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "key_balance": key.balance, - "has_request_body": request_body is not None, - }, - ) - - client = httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, # No timeout - requests can take as long as needed - ) - - try: - # Use the pre-read body if available, otherwise stream - if request_body is not None: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=request_body, - params=prepare_upstream_params(path, request.query_params), - ), - stream=True, - ) - else: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=request.stream(), - params=prepare_upstream_params(path, request.query_params), - ), - stream=True, - ) - - logger.info( - "Received upstream response", - extra={ - "status_code": response.status_code, - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "content_type": response.headers.get("content-type", "unknown"), - }, - ) - - # Map and return errors immediately to provide clear messages - if response.status_code != 200: - try: - mapped_error = await map_upstream_error_response( - request, path, response - ) - finally: - await response.aclose() - await client.aclose() - return mapped_error - - # For chat completions, we need to handle token-based pricing - if path.endswith("chat/completions"): - # Check if client requested streaming - client_wants_streaming = False - if request_body: - try: - request_data = json.loads(request_body) - client_wants_streaming = request_data.get("stream", False) - logger.debug( - "Chat completion request analysis", - extra={ - "client_wants_streaming": client_wants_streaming, - "model": request_data.get("model", "unknown"), - "key_hash": key.hashed_key[:8] + "...", - }, - ) - except json.JSONDecodeError: - logger.warning( - "Failed to parse request body JSON for streaming detection" - ) - - # Handle both streaming and non-streaming responses - content_type = response.headers.get("content-type", "") - upstream_is_streaming = "text/event-stream" in content_type - is_streaming = client_wants_streaming and upstream_is_streaming - - logger.debug( - "Response type analysis", - extra={ - "is_streaming": is_streaming, - "client_wants_streaming": client_wants_streaming, - "upstream_is_streaming": upstream_is_streaming, - "content_type": content_type, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - - if is_streaming and response.status_code == 200: - # Process streaming response and extract cost from the last chunk - result = await handle_streaming_chat_completion( - response, key, max_cost_for_model - ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - result.background = background_tasks - return result - - elif response.status_code == 200: - # Handle non-streaming response - try: - return await handle_non_streaming_chat_completion( - response, key, session, max_cost_for_model - ) - finally: - await response.aclose() - await client.aclose() - - # For all other responses, stream the response - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - - logger.debug( - "Streaming non-chat response", - extra={ - "path": path, - "status_code": response.status_code, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, - ) - - except httpx.RequestError as exc: - await client.aclose() - error_type = type(exc).__name__ - error_details = str(exc) - - logger.error( - "HTTP request error to upstream", - extra={ - "error_type": error_type, - "error_details": error_details, - "method": request.method, - "url": url, - "path": path, - "query_params": dict(request.query_params), - "key_hash": key.hashed_key[:8] + "...", - }, - ) - - # Provide more specific error messages based on the error type - if isinstance(exc, httpx.ConnectError): - error_message = "Unable to connect to upstream service" - elif isinstance(exc, httpx.TimeoutException): - error_message = "Upstream service request timed out" - elif isinstance(exc, httpx.NetworkError): - error_message = "Network error while connecting to upstream service" - else: - error_message = f"Error connecting to upstream service: {error_type}" - - return create_error_response( - "upstream_error", error_message, 502, request=request - ) - - except Exception as exc: - await client.aclose() - tb = traceback.format_exc() - - logger.error( - "Unexpected error in upstream forwarding", - extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, - "path": path, - "query_params": dict(request.query_params), - "key_hash": key.hashed_key[:8] + "...", - "traceback": tb, - }, - ) - - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, - ) +upstream = UpstreamProvider( + base_url=settings.upstream_base_url, + api_key=settings.upstream_api_key, + chat_completions_api_version=settings.chat_completions_api_version, +) @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) @@ -679,7 +101,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, max_cost_for_model) + return await upstream.handle_x_cashu(request, x_cashu, path, max_cost_for_model) elif auth := headers.get("authorization", None): logger.debug( @@ -705,8 +127,10 @@ async def proxy( logger.debug("Processing unauthenticated GET request", extra={"path": path}) # TODO: why is this needed? can we remove it? - headers = prepare_upstream_headers(dict(request.headers)) - return await forward_get_to_upstream(request, path, headers) + headers = upstream.prepare_headers(dict(request.headers)) + return await upstream.forward_get_request( + request, path, headers, map_upstream_error_response + ) # Only pay for request if we have request body data (for completions endpoints) if request_body_dict: @@ -744,11 +168,20 @@ async def proxy( raise # Prepare headers for upstream - headers = prepare_upstream_headers(dict(request.headers)) + headers = upstream.prepare_headers(dict(request.headers)) # Forward to upstream and handle response - response = await forward_to_upstream( - request, path, headers, request_body, key, max_cost_for_model, session + response = await upstream.forward_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + map_upstream_error_response, + handle_streaming_chat_completion, + handle_non_streaming_chat_completion, ) if response.status_code != 200: @@ -850,72 +283,3 @@ async def get_bearer_token_key( }, ) raise - - -async def forward_get_to_upstream( - request: Request, - path: str, - headers: dict, -) -> Response | StreamingResponse: - """Forward request to upstream and handle the response.""" - if path.startswith("v1/"): - path = path.replace("v1/", "") - - url = f"{settings.upstream_base_url}/{path}" - - logger.info( - "Forwarding GET request to upstream", - extra={"url": url, "method": request.method, "path": path}, - ) - - async with httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) as client: - try: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=request.stream(), - params=prepare_upstream_params(path, request.query_params), - ), - ) - - logger.info( - "GET request forwarded successfully", - extra={"path": path, "status_code": response.status_code}, - ) - if response.status_code != 200: - try: - mapped = await map_upstream_error_response(request, path, response) - finally: - await response.aclose() - return mapped - - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - ) - except Exception as exc: - tb = traceback.format_exc() - logger.error( - "Error forwarding GET request", - extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, - "path": path, - "query_params": dict(request.query_params), - "traceback": tb, - }, - ) - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, - ) diff --git a/routstr/upstream.py b/routstr/upstream.py new file mode 100644 index 00000000..23dca318 --- /dev/null +++ b/routstr/upstream.py @@ -0,0 +1,1412 @@ +from __future__ import annotations + +import json +import re +import traceback +from collections.abc import AsyncGenerator, Awaitable, Callable +from typing import TYPE_CHECKING, Mapping + +import httpx + +if TYPE_CHECKING: + from .payment.cost_caculation import CostData, MaxCostData + +from fastapi import BackgroundTasks, HTTPException, Request +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 = "" + ): + self.base_url = base_url + self.api_key = api_key + self.chat_completions_api_version = chat_completions_api_version + + def prepare_headers(self, request_headers: dict) -> dict: + logger.debug( + "Preparing upstream headers", + extra={ + "original_headers_count": len(request_headers), + "has_upstream_api_key": bool(self.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 self.api_key: + headers["Authorization"] = f"Bearer {self.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(self.api_key), + }, + ) + + return headers + + def prepare_params( + self, path: str, query_params: Mapping[str, str] | None + ) -> dict[str, str]: + params: dict[str, str] = dict(query_params or {}) + if path.endswith("chat/completions") and self.chat_completions_api_version: + params["api-version"] = self.chat_completions_api_version + return params + + async def forward_request( + self, + request: Request, + path: str, + headers: dict, + request_body: bytes | None, + 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/", "") + + url = f"{self.base_url}/{path}" + + logger.info( + "Forwarding request to upstream", + extra={ + "url": url, + "method": request.method, + "path": path, + "key_hash": key.hashed_key[:8] + "...", + "key_balance": key.balance, + "has_request_body": request_body is not None, + }, + ) + + client = httpx.AsyncClient( + transport=httpx.AsyncHTTPTransport(retries=1), + timeout=None, + ) + + try: + if request_body is not None: + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=request_body, + params=self.prepare_params(path, request.query_params), + ), + stream=True, + ) + else: + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=request.stream(), + params=self.prepare_params(path, request.query_params), + ), + stream=True, + ) + + logger.info( + "Received upstream response", + extra={ + "status_code": response.status_code, + "path": path, + "key_hash": key.hashed_key[:8] + "...", + "content_type": response.headers.get("content-type", "unknown"), + }, + ) + + if response.status_code != 200: + try: + mapped_error = await map_upstream_error_response( + request, path, response + ) + finally: + await response.aclose() + await client.aclose() + return mapped_error + + if path.endswith("chat/completions"): + client_wants_streaming = False + if request_body: + try: + request_data = json.loads(request_body) + client_wants_streaming = request_data.get("stream", False) + logger.debug( + "Chat completion request analysis", + extra={ + "client_wants_streaming": client_wants_streaming, + "model": request_data.get("model", "unknown"), + "key_hash": key.hashed_key[:8] + "...", + }, + ) + except json.JSONDecodeError: + logger.warning( + "Failed to parse request body JSON for streaming detection" + ) + + content_type = response.headers.get("content-type", "") + upstream_is_streaming = "text/event-stream" in content_type + is_streaming = client_wants_streaming and upstream_is_streaming + + logger.debug( + "Response type analysis", + extra={ + "is_streaming": is_streaming, + "client_wants_streaming": client_wants_streaming, + "upstream_is_streaming": upstream_is_streaming, + "content_type": content_type, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + if is_streaming and response.status_code == 200: + result = await handle_streaming_chat_completion( + response, key, max_cost_for_model + ) + background_tasks = BackgroundTasks() + background_tasks.add_task(response.aclose) + background_tasks.add_task(client.aclose) + result.background = background_tasks + return result + + elif response.status_code == 200: + try: + return await handle_non_streaming_chat_completion( + response, key, session, max_cost_for_model + ) + finally: + await response.aclose() + await client.aclose() + + background_tasks = BackgroundTasks() + background_tasks.add_task(response.aclose) + background_tasks.add_task(client.aclose) + + logger.debug( + "Streaming non-chat response", + extra={ + "path": path, + "status_code": response.status_code, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + return StreamingResponse( + response.aiter_bytes(), + status_code=response.status_code, + headers=dict(response.headers), + background=background_tasks, + ) + + except httpx.RequestError as exc: + await client.aclose() + error_type = type(exc).__name__ + error_details = str(exc) + + logger.error( + "HTTP request error to upstream", + extra={ + "error_type": error_type, + "error_details": error_details, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + if isinstance(exc, httpx.ConnectError): + error_message = "Unable to connect to upstream service" + elif isinstance(exc, httpx.TimeoutException): + error_message = "Upstream service request timed out" + elif isinstance(exc, httpx.NetworkError): + error_message = "Network error while connecting to upstream service" + else: + error_message = f"Error connecting to upstream service: {error_type}" + + return create_error_response( + "upstream_error", error_message, 502, request=request + ) + + except Exception as exc: + await client.aclose() + tb = traceback.format_exc() + + logger.error( + "Unexpected error in upstream forwarding", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "key_hash": key.hashed_key[:8] + "...", + "traceback": tb, + }, + ) + + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) + + async def forward_get_request( + self, + 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/", "") + + url = f"{self.base_url}/{path}" + + logger.info( + "Forwarding GET request to upstream", + extra={"url": url, "method": request.method, "path": path}, + ) + + async with httpx.AsyncClient( + transport=httpx.AsyncHTTPTransport(retries=1), + timeout=None, + ) as client: + try: + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=request.stream(), + params=self.prepare_params(path, request.query_params), + ), + ) + + logger.info( + "GET request forwarded successfully", + extra={"path": path, "status_code": response.status_code}, + ) + if response.status_code != 200: + try: + mapped = await map_upstream_error_response( + request, path, response + ) + finally: + await response.aclose() + return mapped + + return StreamingResponse( + response.aiter_bytes(), + status_code=response.status_code, + headers=dict(response.headers), + ) + except Exception as exc: + tb = traceback.format_exc() + logger.error( + "Error forwarding GET request", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "traceback": tb, + }, + ) + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) + + async def forward_x_cashu_request( + self, + request: Request, + path: str, + headers: dict, + 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/", "") + + url = f"{self.base_url}/{path}" + + logger.debug( + "Forwarding request to upstream", + extra={ + "url": url, + "method": request.method, + "path": path, + "amount": amount, + "unit": unit, + }, + ) + + async with httpx.AsyncClient( + transport=httpx.AsyncHTTPTransport(retries=1), + timeout=None, + ) as client: + try: + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=request.stream(), + params=self.prepare_params(path, request.query_params), + ), + stream=True, + ) + + logger.debug( + "Received upstream response", + extra={ + "status_code": response.status_code, + "path": path, + "response_headers": dict(response.headers), + }, + ) + + if response.status_code != 200: + logger.warning( + "Upstream request failed, processing refund", + extra={ + "status_code": response.status_code, + "path": path, + "amount": amount, + "unit": unit, + }, + ) + + refund_token = await send_refund(amount - 60, unit) + + logger.info( + "Refund processed for failed upstream request", + extra={ + "status_code": response.status_code, + "refund_amount": amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, + }, + ) + + error_response = Response( + content=json.dumps( + { + "error": { + "message": "Error forwarding request to upstream", + "type": "upstream_error", + "code": response.status_code, + "refund_token": refund_token, + } + } + ), + status_code=response.status_code, + media_type="application/json", + ) + error_response.headers["X-Cashu"] = refund_token + return error_response + + if path.endswith("chat/completions"): + logger.debug( + "Processing chat completion response", + extra={"path": path, "amount": amount, "unit": 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 + return result + + background_tasks = BackgroundTasks() + background_tasks.add_task(response.aclose) + background_tasks.add_task(client.aclose) + + logger.debug( + "Streaming non-chat response", + extra={"path": path, "status_code": response.status_code}, + ) + + return StreamingResponse( + response.aiter_bytes(), + status_code=response.status_code, + headers=dict(response.headers), + background=background_tasks, + ) + except Exception as exc: + tb = traceback.format_exc() + logger.error( + "Unexpected error in upstream forwarding", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "traceback": tb, + }, + ) + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) + + async def handle_x_cashu( + self, request: Request, x_cashu_token: str, path: str, max_cost_for_model: int + ) -> Response | StreamingResponse: + from .wallet import recieve_token + + logger.info( + "Processing X-Cashu payment request", + extra={ + "path": path, + "method": request.method, + "token_preview": x_cashu_token[:20] + "..." + if len(x_cashu_token) > 20 + else x_cashu_token, + }, + ) + + try: + headers = dict(request.headers) + amount, unit, mint = await recieve_token(x_cashu_token) + headers = self.prepare_headers(dict(request.headers)) + + logger.info( + "X-Cashu token redeemed successfully", + extra={"amount": amount, "unit": unit, "path": path, "mint": mint}, + ) + + return await self.forward_x_cashu_request( + request, + path, + headers, + amount, + unit, + max_cost_for_model, + handle_x_cashu_chat_completion, + send_refund, + ) + except Exception as e: + error_message = str(e) + logger.error( + "X-Cashu payment request failed", + extra={ + "error": error_message, + "error_type": type(e).__name__, + "path": path, + "method": request.method, + }, + ) + + if "already spent" in error_message.lower(): + return create_error_response( + "token_already_spent", + "The provided CASHU token has already been spent", + 400, + request=request, + token=x_cashu_token, + ) + + if "invalid token" in error_message.lower(): + return create_error_response( + "invalid_token", + "The provided CASHU token is invalid", + 400, + request=request, + token=x_cashu_token, + ) + + if "mint error" in error_message.lower(): + return create_error_response( + "mint_error", + f"CASHU mint error: {error_message}", + 422, + request=request, + token=x_cashu_token, + ) + + return create_error_response( + "cashu_error", + f"CASHU token processing failed: {error_message}", + 400, + 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", + } + }, + ) From 36d55216fedb10e5647e640b6996a851f9f6ffaa Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 4 Oct 2025 17:38:05 +0800 Subject: [PATCH 02/95] fix tests --- tests/integration/test_error_handling_edge_cases.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/integration/test_error_handling_edge_cases.py b/tests/integration/test_error_handling_edge_cases.py index 0d68249c..cb3b3a15 100644 --- a/tests/integration/test_error_handling_edge_cases.py +++ b/tests/integration/test_error_handling_edge_cases.py @@ -48,7 +48,7 @@ class TestNetworkFailureScenarios: ) -> None: """Test proxy behavior when upstream LLM service is down""" # Mock at the routstr level to simulate upstream being down - with patch("routstr.proxy.httpx.AsyncClient") as mock_client_class: + with patch("routstr.upstream.httpx.AsyncClient") as mock_client_class: # Create a mock client instance mock_client = AsyncMock() mock_client_class.return_value = mock_client From c90abe9e791798f93524b3eeeb0a1b6bb51119e4 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 4 Oct 2025 20:12:28 +0800 Subject: [PATCH 03/95] update version --- docs/api/overview.md | 2 +- docs/contributing/code-structure.md | 2 +- docs/getting-started/quickstart.md | 2 +- pyproject.toml | 2 +- routstr/core/main.py | 2 +- uv.lock | 2 +- 6 files changed, 6 insertions(+), 6 deletions(-) diff --git a/docs/api/overview.md b/docs/api/overview.md index 16d8d77f..cb328382 100644 --- a/docs/api/overview.md +++ b/docs/api/overview.md @@ -347,7 +347,7 @@ GET /health Response: { "status": "healthy", - "version": "0.1.4", + "version": "0.1.5", "timestamp": "2024-01-01T00:00:00Z", "checks": { "database": "ok", diff --git a/docs/contributing/code-structure.md b/docs/contributing/code-structure.md index be2df5d5..97f6d716 100644 --- a/docs/contributing/code-structure.md +++ b/docs/contributing/code-structure.md @@ -348,7 +348,7 @@ Project metadata and dependencies: ```toml [project] name = "routstr" -version = "0.1.4" +version = "0.1.5" dependencies = [ "fastapi[standard]>=0.115", "sqlmodel>=0.0.24", diff --git a/docs/getting-started/quickstart.md b/docs/getting-started/quickstart.md index ba5d6751..0d46ebc6 100644 --- a/docs/getting-started/quickstart.md +++ b/docs/getting-started/quickstart.md @@ -67,7 +67,7 @@ You should see: { "name": "ARoutstrNode", "description": "A Routstr Node", - "version": "0.1.4", + "version": "0.1.5", "npub": "", "mints": ["https://mint.minibits.cash/Bitcoin"], "models": {...} diff --git a/pyproject.toml b/pyproject.toml index 8d5843b1..b636bb0e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "routstr" -version = "0.1.4" +version = "0.1.5" description = "Payment proxy for your LLM endpoint using cashu and nostr." readme = "README.md" requires-python = ">=3.11" diff --git a/routstr/core/main.py b/routstr/core/main.py index 1669a7bd..b0fbb22f 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -30,7 +30,7 @@ from .settings import settings as global_settings setup_logging() logger = get_logger(__name__) -__version__ = "0.1.4-dev" +__version__ = "0.1.5-dev" @asynccontextmanager diff --git a/uv.lock b/uv.lock index f8a0ed42..7058cf59 100644 --- a/uv.lock +++ b/uv.lock @@ -1783,7 +1783,7 @@ wheels = [ [[package]] name = "routstr" -version = "0.1.4" +version = "0.1.5" source = { editable = "." } dependencies = [ { name = "aiosqlite" }, From 719c091145b153dcf90883eb42d461e8d31a7322 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 4 Oct 2025 20:19:04 +0800 Subject: [PATCH 04/95] ruff fmt --- routstr/wallet.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index c2bd83d0..fe34d7bb 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -152,9 +152,7 @@ async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wal global _wallets id = f"{mint_url}_{unit}" if id not in _wallets: - _wallets[id] = await Wallet.with_db( - mint_url, db=".wallet", unit=unit - ) + _wallets[id] = await Wallet.with_db(mint_url, db=".wallet", unit=unit) if load: await _wallets[id].load_mint() From 384a149600a17ef1c8a7084da476deaace18746e Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 4 Oct 2025 20:52:15 +0800 Subject: [PATCH 05/95] fix tests --- tests/integration/conftest.py | 2 ++ .../test_general_info_endpoints.py | 25 +++++-------------- 2 files changed, 8 insertions(+), 19 deletions(-) diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 311214f6..220b3c9c 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -63,6 +63,8 @@ else: # Set test environment variables before importing the app os.environ.update(test_env) +os.environ.pop("ADMIN_PASSWORD", None) + from routstr.core.db import ApiKey, get_session # noqa: E402 from routstr.core.main import app, lifespan # noqa: E402 diff --git a/tests/integration/test_general_info_endpoints.py b/tests/integration/test_general_info_endpoints.py index f9d52c25..cfde36e1 100644 --- a/tests/integration/test_general_info_endpoints.py +++ b/tests/integration/test_general_info_endpoints.py @@ -271,36 +271,23 @@ async def test_models_endpoint_accept_headers(integration_client: AsyncClient) - async def test_admin_endpoint_unauthenticated( integration_client: AsyncClient, db_snapshot: Any ) -> None: - """Test GET /admin/ endpoint without authentication""" - - # Capture initial database state + """Test GET /admin/ endpoint without authentication shows setup form""" await db_snapshot.capture() response = await integration_client.get("/admin/") - # Should return 200 with login form (not 401/403) assert response.status_code == 200 assert "text/html" in response.headers["content-type"] - # Response should be HTML html_content = response.text assert "" in html_content assert "" in html_content + assert "" in html_content or " +""" + + +def upstream_providers_page() -> str: + return ( + f""" + + + + {UPSTREAM_PROVIDERS_JS} + + """ + + """ + + ← Back to Dashboard +

Upstream Providers

+ +
+

Providers

+
+ +
+ + + + + + + + + + + + + +
IDTypeBase URLStatusActions
Loading…
+
+ + + + + + + + + """ + ) + + +@admin_router.get("/upstream-providers", response_class=HTMLResponse) +async def admin_upstream_providers(request: Request) -> str: + if is_admin_authenticated(request): + return upstream_providers_page() + return admin_auth() + + @admin_router.get("/api/models", dependencies=[Depends(require_admin_api)]) async def get_models_admin_api(request: Request) -> list[dict[str, object]]: items = await list_models() return [m.dict() for m in items] # type: ignore +class ModelCreate(BaseModel): + id: str + name: str + description: str + created: int + context_length: int + architecture: dict[str, object] + pricing: dict[str, object] + per_request_limits: dict[str, object] | None = None + top_provider: dict[str, object] | None = None + upstream_provider_id: int | None = None + enabled: bool = True + + @admin_router.post("/api/models", dependencies=[Depends(require_admin_api)]) -async def create_model_admin_api(payload: Model) -> dict[str, object]: +async def create_model_admin_api(payload: ModelCreate) -> dict[str, object]: async with create_session() as session: exists = await session.get(ModelRow, payload.id) if exists: raise HTTPException( status_code=409, detail="Model with this ID already exists" ) - pricing_dict = payload.pricing.dict() - for k in ("max_prompt_cost", "max_completion_cost", "max_cost"): - pricing_dict.pop(k, None) row = ModelRow( id=payload.id, name=payload.name, description=payload.description, created=int(payload.created), context_length=int(payload.context_length), - architecture=json.dumps(payload.architecture.dict()), - pricing=json.dumps(pricing_dict), + architecture=json.dumps(payload.architecture), + pricing=json.dumps(payload.pricing), sats_pricing=None, per_request_limits=( json.dumps(payload.per_request_limits) @@ -1416,16 +2121,17 @@ async def create_model_admin_api(payload: Model) -> dict[str, object]: else None ), top_provider=( - json.dumps(payload.top_provider.dict()) - if payload.top_provider - else None + json.dumps(payload.top_provider) if payload.top_provider else None ), + upstream_provider_id=payload.upstream_provider_id, + enabled=payload.enabled, ) session.add(row) await session.commit() + await session.refresh(row) - created_model = await get_model_by_id(payload.id) - return created_model.dict() if created_model else {"id": payload.id} # type: ignore + await refresh_model_maps() + return _row_to_model(row).dict() # type: ignore @admin_router.post("/api/models/batch", dependencies=[Depends(require_admin_api)]) @@ -1475,6 +2181,8 @@ async def batch_create_models(payload: dict[str, object]) -> dict[str, int]: created += 1 if created: await session.commit() + if created: + await refresh_model_maps() return {"created": created, "skipped": skipped} @@ -1482,16 +2190,33 @@ async def batch_create_models(payload: dict[str, object]) -> dict[str, int]: "/api/models/{model_id:path}", dependencies=[Depends(require_admin_api)] ) async def get_model_admin_api(model_id: str) -> dict[str, object]: - model = await get_model_by_id(model_id) - if not model: - raise HTTPException(status_code=404, detail="Model not found") - return model.dict() # type: ignore + async with create_session() as session: + row = await session.get(ModelRow, model_id) + if not row: + raise HTTPException(status_code=404, detail="Model not found") + return _row_to_model(row).dict() # type: ignore + + +class ModelUpdate(BaseModel): + id: str + name: str + description: str + created: int + context_length: int + architecture: dict[str, object] + pricing: dict[str, object] + per_request_limits: dict[str, object] | None = None + top_provider: dict[str, object] | None = None + upstream_provider_id: int | None = None + enabled: bool = True @admin_router.patch( "/api/models/{model_id:path}", dependencies=[Depends(require_admin_api)] ) -async def update_model_admin_api(model_id: str, payload: Model) -> dict[str, object]: +async def update_model_admin_api( + model_id: str, payload: ModelUpdate +) -> dict[str, object]: if payload.id != model_id: raise HTTPException(status_code=400, detail="Path id does not match payload id") @@ -1504,11 +2229,8 @@ async def update_model_admin_api(model_id: str, payload: Model) -> dict[str, obj row.description = payload.description row.created = int(payload.created) row.context_length = int(payload.context_length) - row.architecture = json.dumps(payload.architecture.dict()) - pricing_dict = payload.pricing.dict() - for k in ("max_prompt_cost", "max_completion_cost", "max_cost"): - pricing_dict.pop(k, None) - row.pricing = json.dumps(pricing_dict) + row.architecture = json.dumps(payload.architecture) + row.pricing = json.dumps(payload.pricing) row.sats_pricing = None row.per_request_limits = ( json.dumps(payload.per_request_limits) @@ -1516,16 +2238,17 @@ async def update_model_admin_api(model_id: str, payload: Model) -> dict[str, obj else None ) row.top_provider = ( - json.dumps(payload.top_provider.dict()) if payload.top_provider else None + json.dumps(payload.top_provider) if payload.top_provider else None ) + row.upstream_provider_id = payload.upstream_provider_id + row.enabled = payload.enabled session.add(row) await session.commit() + await session.refresh(row) - updated = await get_model_by_id(model_id) - if not updated: - raise HTTPException(status_code=404, detail="Model not found after update") - return updated.dict() # type: ignore + await refresh_model_maps() + return _row_to_model(row).dict() # type: ignore @admin_router.delete( @@ -1538,6 +2261,7 @@ async def delete_model_admin_api(model_id: str) -> dict[str, object]: raise HTTPException(status_code=404, detail="Model not found") await session.delete(row) await session.commit() + await refresh_model_maps() return {"ok": True, "deleted_id": model_id} @@ -1549,9 +2273,188 @@ async def delete_all_models_admin_api() -> dict[str, object]: for row in rows: await session.delete(row) # type: ignore await session.commit() + await refresh_model_maps() return {"ok": True, "deleted": "all"} +class UpstreamProviderCreate(BaseModel): + provider_type: str + base_url: str + api_key: str + api_version: str | None = None + enabled: bool = True + + +class UpstreamProviderUpdate(BaseModel): + provider_type: str | None = None + base_url: str | None = None + api_key: str | None = None + api_version: str | None = None + enabled: bool | None = None + + +@admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) +async def get_upstream_providers() -> list[dict[str, object]]: + async with create_session() as session: + result = await session.exec(select(UpstreamProviderRow)) + providers = result.all() + return [ + { + "id": p.id, + "provider_type": p.provider_type, + "base_url": p.base_url, + "api_key": "[REDACTED]" if p.api_key else "", + "api_version": p.api_version, + "enabled": p.enabled, + } + for p in providers + ] + + +@admin_router.post("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) +async def create_upstream_provider( + payload: UpstreamProviderCreate, +) -> dict[str, object]: + async with create_session() as session: + result = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.base_url == payload.base_url + ) + ) + if result.first(): + raise HTTPException( + status_code=409, detail="Provider with this base URL already exists" + ) + + provider = UpstreamProviderRow( + provider_type=payload.provider_type, + base_url=payload.base_url, + api_key=payload.api_key, + api_version=payload.api_version, + enabled=payload.enabled, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + + await reinitialize_upstreams() + return { + "id": provider.id, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + "api_key": "[REDACTED]", + "api_version": provider.api_version, + "enabled": provider.enabled, + } + + +@admin_router.get( + "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] +) +async def get_upstream_provider(provider_id: int) -> dict[str, object]: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + return { + "id": provider.id, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + "api_key": "[REDACTED]" if provider.api_key else "", + "api_version": provider.api_version, + "enabled": provider.enabled, + } + + +@admin_router.patch( + "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] +) +async def update_upstream_provider( + provider_id: int, payload: UpstreamProviderUpdate +) -> dict[str, object]: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + + if payload.provider_type is not None: + provider.provider_type = payload.provider_type + if payload.base_url is not None: + provider.base_url = payload.base_url + if payload.api_key is not None: + provider.api_key = payload.api_key + if payload.api_version is not None: + provider.api_version = payload.api_version + if payload.enabled is not None: + provider.enabled = payload.enabled + + session.add(provider) + await session.commit() + await session.refresh(provider) + + await reinitialize_upstreams() + return { + "id": provider.id, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + "api_key": "[REDACTED]", + "api_version": provider.api_version, + "enabled": provider.enabled, + } + + +@admin_router.delete( + "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] +) +async def delete_upstream_provider(provider_id: int) -> dict[str, object]: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + await session.delete(provider) + await session.commit() + await reinitialize_upstreams() + return {"ok": True, "deleted_id": provider_id} + + +@admin_router.get( + "/api/upstream-providers/{provider_id}/models", + dependencies=[Depends(require_admin_api)], +) +async def get_provider_models(provider_id: int) -> dict[str, object]: + from ..upstream import _instantiate_provider + + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + + db_models = await list_models( + session=session, upstream_id=provider_id, include_disabled=True + ) + + remote_models = [] + upstream_instance = _instantiate_provider(provider) + if upstream_instance: + try: + models = await upstream_instance.fetch_models() + remote_models = [m.dict() for m in models] + except Exception as e: + logger.error( + f"Failed to fetch models from {provider.provider_type}: {e}" + ) + + return { + "provider": { + "id": provider.id, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + }, + "db_models": [m.dict() for m in db_models], + "remote_models": remote_models, + } + + DASHBOARD_CSS: str = """ * { margin: 0; padding: 0; box-sizing: border-box; } body { font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif; background: #f5f7fa; color: #2c3e50; line-height: 1.6; padding: 2rem; } @@ -1591,8 +2494,8 @@ button:disabled { background: #a0aec0; cursor: not-allowed; transform: none; } @keyframes slideIn { from { transform: translateY(-20px); opacity: 0; } to { transform: translateY(0); opacity: 1; } } .close { color: #a0aec0; float: right; font-size: 28px; font-weight: bold; cursor: pointer; margin: -10px -10px 0 0; } .close:hover { color: #2d3748; } -input[type="number"], input[type="text"], select { width: 100%; padding: 10px; margin: 10px 0; border: 2px solid #e2e8f0; border-radius: 6px; font-size: 16px; transition: border 0.2s; } -input[type="number"]:focus, input[type="text"]:focus, select:focus { outline: none; border-color: #4299e1; } +input[type="number"], input[type="text"], input[type="password"], select { width: 100%; padding: 10px; margin: 10px 0; border: 2px solid #e2e8f0; border-radius: 6px; font-size: 16px; transition: border 0.2s; } +input[type="number"]:focus, input[type="text"]:focus, input[type="password"]:focus, select:focus { outline: none; border-color: #4299e1; } .warning { color: #e53e3e; font-weight: 600; margin: 10px 0; padding: 10px; background: #fff5f5; border-radius: 6px; } """ diff --git a/routstr/core/db.py b/routstr/core/db.py index 9f886791..744d5ed1 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -5,7 +5,7 @@ from typing import AsyncGenerator from alembic import command from alembic.config import Config from sqlalchemy.ext.asyncio.engine import create_async_engine -from sqlmodel import Field, SQLModel, func, select +from sqlmodel import Field, Relationship, SQLModel, func, select from sqlmodel.ext.asyncio.session import AsyncSession from .logging import get_logger @@ -64,6 +64,26 @@ class ModelRow(SQLModel, table=True): # type: ignore sats_pricing: str | None = Field(default=None) per_request_limits: str | None = Field(default=None) top_provider: str | None = Field(default=None) + enabled: bool = Field(default=True, description="Whether this model is enabled") + upstream_provider_id: int | None = Field( + default=None, foreign_key="upstream_providers.id" + ) + upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models") + + +class UpstreamProviderRow(SQLModel, table=True): # type: ignore + __tablename__ = "upstream_providers" + id: int | None = Field(default=None, primary_key=True) + provider_type: str = Field( + description="Provider type: generic, openai, azure, openrouter" + ) + base_url: str = Field(unique=True, description="Base URL of the upstream API") + api_key: str = Field(description="API key for the upstream provider") + api_version: str | None = Field( + default=None, description="API version for Azure OpenAI" + ) + enabled: bool = Field(default=True, description="Whether this provider is enabled") + models: list["ModelRow"] = Relationship(back_populates="upstream_provider") async def balances_for_mint_and_unit( diff --git a/routstr/core/main.py b/routstr/core/main.py index 82a7a210..b0266743 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -11,12 +11,10 @@ from ..balance import balance_router, deprecated_wallet_router from ..discovery import providers_cache_refresher, providers_router from ..nip91 import announce_provider from ..payment.models import ( - ensure_models_bootstrapped, models_router, - refresh_models_periodically, update_sats_pricing, ) -from ..proxy import proxy_router +from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..wallet import periodic_payout from .admin import admin_router from .db import create_session, init_db, run_migrations @@ -42,6 +40,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: nip91_task = None providers_task = None models_refresh_task = None + model_maps_refresh_task = None try: # Run database migrations on startup @@ -65,10 +64,18 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: except Exception: pass - await ensure_models_bootstrapped() + # await ensure_models_bootstrapped() + await initialize_upstreams() + + from ..proxy import get_upstreams + from ..upstream import refresh_upstreams_models_periodically + pricing_task = asyncio.create_task(update_sats_pricing()) if global_settings.models_refresh_interval_seconds > 0: - models_refresh_task = asyncio.create_task(refresh_models_periodically()) + models_refresh_task = asyncio.create_task( + refresh_upstreams_models_periodically(get_upstreams()) + ) + model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically()) payout_task = asyncio.create_task(periodic_payout()) nip91_task = asyncio.create_task(announce_provider()) providers_task = asyncio.create_task(providers_cache_refresher()) @@ -94,6 +101,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: providers_task.cancel() if models_refresh_task is not None: models_refresh_task.cancel() + if model_maps_refresh_task is not None: + model_maps_refresh_task.cancel() try: tasks_to_wait = [] @@ -107,6 +116,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(providers_task) if models_refresh_task is not None: tasks_to_wait.append(models_refresh_task) + if model_maps_refresh_task is not None: + tasks_to_wait.append(model_maps_refresh_task) if tasks_to_wait: await asyncio.gather(*tasks_to_wait, return_exceptions=True) diff --git a/routstr/payment/cost_caculation.py b/routstr/payment/cost_caculation.py index 50b82e53..bf4b42e9 100644 --- a/routstr/payment/cost_caculation.py +++ b/routstr/payment/cost_caculation.py @@ -1,12 +1,8 @@ -import json import math from pydantic.v1 import BaseModel -from sqlmodel import select -from sqlmodel.ext.asyncio.session import AsyncSession from ..core import get_logger -from ..core.db import ModelRow from ..core.settings import settings logger = get_logger(__name__) @@ -29,7 +25,7 @@ class CostDataError(BaseModel): async def calculate_cost( - response_data: dict, max_cost: int, session: AsyncSession | None = None + response_data: dict, max_cost: int, session: object | None = None ) -> CostData | MaxCostData | CostDataError: """ Calculate the cost of an API request based on token usage. @@ -74,18 +70,20 @@ async def calculate_cost( float(settings.fixed_per_1k_output_tokens) * 1000.0 ) - if not settings.fixed_pricing and session is not None: + if not settings.fixed_pricing: response_model = response_data.get("model", "") logger.debug( "Using model-based pricing", extra={"model": response_model}, ) - result = await session.exec(select(ModelRow.id)) # type: ignore - available_ids = [ - row[0] if isinstance(row, tuple) else row for row in result.all() - ] - if response_model not in available_ids: + from ..proxy import get_upstreams + from ..upstream import get_model_with_override + + upstreams = get_upstreams() + model_obj = await get_model_with_override(response_model, upstreams) + + if not model_obj: logger.error( "Invalid model in response", extra={"response_model": response_model}, @@ -95,8 +93,7 @@ async def calculate_cost( code="model_not_found", ) - row = await session.get(ModelRow, response_model) - if row is None or not row.sats_pricing: + if not model_obj.sats_pricing: logger.error( "Model pricing not defined", extra={"model": response_model, "model_id": response_model}, @@ -106,9 +103,8 @@ async def calculate_cost( ) try: - sats_pricing = json.loads(row.sats_pricing) - mspp = float(sats_pricing.get("prompt", 0)) - mspc = float(sats_pricing.get("completion", 0)) + mspp = float(model_obj.sats_pricing.prompt) + mspc = float(model_obj.sats_pricing.completion) except Exception: return CostDataError(message="Invalid pricing data", code="pricing_invalid") diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 37c7c8be..ef879b0a 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -1,13 +1,12 @@ import json import math +from typing import Any from fastapi import HTTPException, Response from fastapi.requests import Request -from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession from ..core import get_logger -from ..core.db import ModelRow from ..core.settings import settings from ..wallet import deserialize_token_from_string from .models import Pricing @@ -84,19 +83,19 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N async def get_max_cost_for_model( - model: str, session: AsyncSession | None = None + model: str, + session: AsyncSession | None = None, + model_obj: Any | None = None, ) -> int: - """Get the maximum cost for a specific model.""" + """Get the maximum cost for a specific model from providers with overrides.""" logger.debug( "Getting max cost for model", extra={ "model": model, "fixed_pricing": settings.fixed_pricing, - "has_models": True, }, ) - # Fixed pricing: always use fixed_cost_per_request if settings.fixed_pricing: default_cost_msats = settings.fixed_cost_per_request * 1000 logger.debug( @@ -105,43 +104,42 @@ async def get_max_cost_for_model( ) return max(settings.min_request_msat, default_cost_msats) - if session is None: - # Without a DB session, we can't resolve model pricing; fall back to fixed cost - fallback_msats = settings.fixed_cost_per_request * 1000 - logger.warning( - "No DB session provided for model pricing; using fixed cost", - extra={"requested_model": model, "using_default_cost": fallback_msats}, - ) - return max(settings.min_request_msat, fallback_msats) + if not model_obj: + from ..proxy import get_upstreams + from ..upstream import get_model_with_override - result = await session.exec(select(ModelRow.id)) # type: ignore - available_ids = [row[0] if isinstance(row, tuple) else row for row in result.all()] - if model not in available_ids: - # If no models or unknown model, fall back to fixed cost if provided, else minimal default + upstreams = get_upstreams() + model_obj = await get_model_with_override(model, upstreams) + + if not model_obj: fallback_msats = settings.fixed_cost_per_request * 1000 logger.warning( - "Model not found in available models", + "Model not found in providers or overrides", extra={ "requested_model": model, - "available_models": available_ids, "using_default_cost": fallback_msats, }, ) return max(settings.min_request_msat, fallback_msats) - row = await session.get(ModelRow, model) - if row and row.sats_pricing: + if model_obj.sats_pricing: try: - sats = Pricing(**json.loads(row.sats_pricing)) # type: ignore - max_cost = sats.max_cost * 1000 * (1 - settings.tolerance_percentage / 100) + max_cost = ( + model_obj.sats_pricing.max_cost + * 1000 + * (1 - settings.tolerance_percentage / 100) + ) logger.debug( "Found model-specific max cost", extra={"model": model, "max_cost_msats": max_cost}, ) calculated_msats = int(max_cost) return max(settings.min_request_msat, calculated_msats) - except Exception: - pass + except Exception as e: + logger.error( + "Error calculating max cost from model pricing", + extra={"model": model, "error": str(e)}, + ) logger.warning( "Model pricing not found, using fixed cost", @@ -220,16 +218,19 @@ def estimate_tokens(messages: list) -> int: async def get_model_cost_info( model_id: str, session: AsyncSession | None = None ) -> Pricing | None: + """Get model pricing info from providers with database overrides.""" if not model_id or model_id == "unknown": return None - if session is None: - return None - row = await session.get(ModelRow, model_id) - if row and row.sats_pricing: - try: - return Pricing(**json.loads(row.sats_pricing)) # type: ignore - except Exception: - return None + + from ..proxy import get_upstreams + from ..upstream import get_model_with_override + + upstreams = get_upstreams() + model_obj = await get_model_with_override(model_id, upstreams) + + if model_obj and model_obj.sats_pricing: + return model_obj.sats_pricing + return None diff --git a/routstr/payment/models.py b/routstr/payment/models.py index c064a8df..a5b0755e 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -4,6 +4,7 @@ import random from pathlib import Path from urllib.request import urlopen +import httpx from fastapi import APIRouter, Depends from pydantic.v1 import BaseModel from sqlmodel import select @@ -56,6 +57,12 @@ class Model(BaseModel): sats_pricing: Pricing | None = None per_request_limits: dict | None = None top_provider: TopProvider | None = None + enabled: bool = True + upstream_provider_id: int | None = None + canonical_slug: str | None = None + + def __hash__(self) -> int: + return hash(self.id) def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: @@ -97,6 +104,47 @@ def fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: return [] +async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: + """Asynchronously fetch model information from OpenRouter API.""" + base_url = "https://openrouter.ai/api/v1" + + try: + async with httpx.AsyncClient() as client: + response = await client.get(f"{base_url}/models", timeout=30) + response.raise_for_status() + data = response.json() + + models_data: list[dict] = [] + for model in data.get("data", []): + model_id = model.get("id", "") + + if source_filter: + source_prefix = f"{source_filter}/" + if not model_id.startswith(source_prefix): + continue + + model = dict(model) + model["id"] = model_id[len(source_prefix) :] + model_id = model["id"] + + if ( + "(free)" in model.get("name", "") + or model_id == "openrouter/auto" + or model_id == "google/gemini-2.5-pro-exp-03-25" + or model_id == "opengvlab/internvl3-78b" + or model_id == "openrouter/sonoma-dusk-alpha" + or model_id == "openrouter/sonoma-sky-alpha" + ): + continue + + models_data.append(model) + + return models_data + except Exception as e: + logger.error(f"Error (async) fetching models from OpenRouter API: {e}") + return [] + + def is_openrouter_upstream() -> bool: try: base = (settings.upstream_base_url or "").strip().rstrip("/") @@ -188,10 +236,13 @@ def _row_to_model(row: ModelRow) -> Model: sats_pricing=Pricing.parse_obj(sats_pricing) if sats_pricing else None, per_request_limits=per_request_limits, top_provider=TopProvider.parse_obj(top_provider) if top_provider else None, + enabled=row.enabled, + upstream_provider_id=row.upstream_provider_id, + canonical_slug=getattr(row, "canonical_slug", None), ) -def _model_to_row_payload(model: Model) -> dict[str, str | int | None]: +def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]: return { "id": model.id, "name": model.name, @@ -209,18 +260,28 @@ def _model_to_row_payload(model: Model) -> dict[str, str | int | None]: "top_provider": json.dumps(model.top_provider.dict()) if model.top_provider is not None else None, + "enabled": model.enabled, + "upstream_provider_id": model.upstream_provider_id, } -async def list_models(session: AsyncSession | None = None) -> list[Model]: +async def list_models( + session: AsyncSession | None = None, + upstream_id: int | None = None, + include_disabled: bool = False, +) -> list[Model]: + from sqlmodel import select + + query = select(ModelRow) + if upstream_id is not None: + query = query.where(ModelRow.upstream_provider_id == upstream_id) + if not include_disabled: + query = query.where(ModelRow.enabled) + if session is not None: - result = await session.exec(select(ModelRow)) # type: ignore - rows = result.all() - return [_row_to_model(r) for r in rows] + return [_row_to_model(r) for r in (await session.exec(query)).all()] # type: ignore async with create_session() as s: - result = await s.exec(select(ModelRow)) # type: ignore - rows = result.all() - return [_row_to_model(r) for r in rows] + return [_row_to_model(r) for r in (await s.exec(query)).all()] # type: ignore async def get_model_by_id( @@ -228,10 +289,101 @@ async def get_model_by_id( ) -> Model | None: if session is not None: row = await session.get(ModelRow, model_id) - return _row_to_model(row) if row else None + return _row_to_model(row) if row and row.enabled else None async with create_session() as s: row = await s.get(ModelRow, model_id) - return _row_to_model(row) if row else None + return _row_to_model(row) if row and row.enabled else None + + +def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model: + """Update a model's sats_pricing based on USD pricing and exchange rate. + + Args: + model: Model object to update + sats_to_usd: Current sats to USD exchange rate + + Returns: + Updated Model object with new sats_pricing + """ + try: + sats = Pricing.parse_obj( + {k: v / sats_to_usd for k, v in model.pricing.dict().items()} + ) + + min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1))) + min_req_sats = float(min_req_msat) / 1000.0 + if sats.request <= 0.0: + sats.request = min_req_sats + + mspp = sats.prompt + mspc = sats.completion + + if model.top_provider and ( + model.top_provider.context_length + or model.top_provider.max_completion_tokens + ): + if (cl := model.top_provider.context_length) and ( + mct := model.top_provider.max_completion_tokens + ): + max_prompt_cost = (cl - mct) * mspp + max_completion_cost = mct * mspc + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + elif cl := model.top_provider.context_length: + max_prompt_cost = cl * 0.8 * mspp + max_completion_cost = cl * 0.2 * mspc + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + elif mct := model.top_provider.max_completion_tokens: + max_prompt_cost = mct * 4 * mspp + max_completion_cost = mct * mspc + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + elif model.context_length: + max_prompt_cost = mspp * model.context_length * 0.8 + max_completion_cost = mspc * model.context_length * 0.2 + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + else: + p = mspp * 1_000_000 + c = mspc * 32_000 + r = sats.request * 100_000 + i = sats.image * 100 + w = sats.web_search * 1000 + ir = sats.internal_reasoning * 100 + sats.max_prompt_cost = p + sats.max_completion_cost = c + sats.max_cost = p + c + r + i + w + ir + + if (sats.max_cost or 0.0) < min_req_sats: + sats.max_cost = min_req_sats + + return Model( + id=model.id, + name=model.name, + created=model.created, + description=model.description, + context_length=model.context_length, + architecture=model.architecture, + pricing=model.pricing, + sats_pricing=sats, + per_request_limits=model.per_request_limits, + top_provider=model.top_provider, + ) + except Exception as e: + logger.error( + "Failed to update sats pricing for model", + extra={ + "model_id": model.id, + "error": str(e), + "error_type": type(e).__name__, + }, + ) + return model async def ensure_models_bootstrapped() -> None: @@ -285,113 +437,134 @@ async def ensure_models_bootstrapped() -> None: await s.commit() -async def update_sats_pricing() -> None: - while True: - try: +async def _update_sats_pricing_once() -> None: + """Update sats pricing once for all provider models and database overrides.""" + from ..proxy import get_upstreams + + sats_to_usd = await sats_usd_ask_price() + upstreams = get_upstreams() + + updated_count = 0 + + for upstream in upstreams: + updated_models = [ + _update_model_sats_pricing(m, sats_to_usd) + for m in upstream.get_cached_models() + ] + upstream._models_cache = updated_models + upstream._models_by_id = {m.id: m for m in updated_models} + updated_count += len(updated_models) + + async with create_session() as s: + result = await s.exec( + select(ModelRow).where(ModelRow.upstream_provider_id.isnot(None)) # type: ignore + ) # type: ignore + rows = result.all() + changed = 0 + for row in rows: try: - if not settings.enable_pricing_refresh: - return - except Exception: - pass - sats_to_usd = await sats_usd_ask_price() - async with create_session() as s: - result = await s.exec(select(ModelRow)) # type: ignore - rows = result.all() - changed = 0 - for row in rows: - try: - pricing = Pricing.parse_obj(json.loads(row.pricing)) - top_provider = ( - TopProvider.parse_obj(json.loads(row.top_provider)) - if row.top_provider - else None - ) - sats = Pricing.parse_obj( - {k: v / sats_to_usd for k, v in pricing.dict().items()} - ) - # Enforce minimum per-request charge floor in sats - try: - min_req_msat = max( - 1, int(getattr(settings, "min_request_msat", 1)) - ) - except Exception: - min_req_msat = 1 - min_req_sats = float(min_req_msat) / 1000.0 - if sats.request <= 0.0: - sats.request = min_req_sats - mspp = sats.prompt - mspc = sats.completion - if top_provider and ( - top_provider.context_length - or top_provider.max_completion_tokens - ): - if (cl := top_provider.context_length) and ( - mct := top_provider.max_completion_tokens - ): - max_prompt_cost = (cl - mct) * mspp - max_completion_cost = mct * mspc - sats.max_prompt_cost = max_prompt_cost - sats.max_completion_cost = max_completion_cost - sats.max_cost = max_prompt_cost + max_completion_cost - elif cl := top_provider.context_length: - max_prompt_cost = cl * 0.8 * mspp - max_completion_cost = cl * 0.2 * mspc - sats.max_prompt_cost = max_prompt_cost - sats.max_completion_cost = max_completion_cost - sats.max_cost = max_prompt_cost + max_completion_cost - elif mct := top_provider.max_completion_tokens: - max_prompt_cost = mct * 4 * mspp - max_completion_cost = mct * mspc - sats.max_prompt_cost = max_prompt_cost - sats.max_completion_cost = max_completion_cost - sats.max_cost = max_prompt_cost + max_completion_cost - else: - max_prompt_cost = 1_000_000 * mspp - max_completion_cost = 32_000 * mspc - sats.max_prompt_cost = max_prompt_cost - sats.max_completion_cost = max_completion_cost - sats.max_cost = max_prompt_cost + max_completion_cost - elif row.context_length: - max_prompt_cost = mspp * row.context_length * 0.8 - max_completion_cost = mspc * row.context_length * 0.2 - sats.max_prompt_cost = max_prompt_cost - sats.max_completion_cost = max_completion_cost - sats.max_cost = max_prompt_cost + max_completion_cost - else: - p = mspp * 1_000_000 - c = mspc * 32_000 - r = sats.request * 100_000 - i = sats.image * 100 - w = sats.web_search * 1000 - ir = sats.internal_reasoning * 100 - sats.max_prompt_cost = p - sats.max_completion_cost = c - sats.max_cost = p + c + r + i + w + ir + pricing = Pricing.parse_obj(json.loads(row.pricing)) + top_provider = ( + TopProvider.parse_obj(json.loads(row.top_provider)) + if row.top_provider + else None + ) + sats = Pricing.parse_obj( + {k: v / sats_to_usd for k, v in pricing.dict().items()} + ) + min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1))) + min_req_sats = float(min_req_msat) / 1000.0 + if sats.request <= 0.0: + sats.request = min_req_sats + mspp = sats.prompt + mspc = sats.completion + if top_provider and ( + top_provider.context_length or top_provider.max_completion_tokens + ): + if (cl := top_provider.context_length) and ( + mct := top_provider.max_completion_tokens + ): + max_prompt_cost = (cl - mct) * mspp + max_completion_cost = mct * mspc + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + elif cl := top_provider.context_length: + max_prompt_cost = cl * 0.8 * mspp + max_completion_cost = cl * 0.2 * mspc + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + elif mct := top_provider.max_completion_tokens: + max_prompt_cost = mct * 4 * mspp + max_completion_cost = mct * mspc + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + else: + max_prompt_cost = 1_000_000 * mspp + max_completion_cost = 32_000 * mspc + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + elif row.context_length: + max_prompt_cost = mspp * row.context_length * 0.8 + max_completion_cost = mspc * row.context_length * 0.2 + sats.max_prompt_cost = max_prompt_cost + sats.max_completion_cost = max_completion_cost + sats.max_cost = max_prompt_cost + max_completion_cost + else: + p = mspp * 1_000_000 + c = mspc * 32_000 + r = sats.request * 100_000 + i = sats.image * 100 + w = sats.web_search * 1000 + ir = sats.internal_reasoning * 100 + sats.max_prompt_cost = p + sats.max_completion_cost = c + sats.max_cost = p + c + r + i + w + ir - # Ensure overall minimum per-request total cost floor - if (sats.max_cost or 0.0) < min_req_sats: - sats.max_cost = min_req_sats + if (sats.max_cost or 0.0) < min_req_sats: + sats.max_cost = min_req_sats - new_json = json.dumps(sats.dict()) - if row.sats_pricing != new_json: - row.sats_pricing = new_json - s.add(row) - changed += 1 - except Exception as per_row_error: - logger.error( - "Failed to update pricing for model", - extra={ - "model_id": row.id, - "error": str(per_row_error), - "error_type": type(per_row_error).__name__, - }, - ) - if changed: - await s.commit() - except asyncio.CancelledError: - break - except Exception as e: - logger.error(f"Error updating sats pricing: {e}") + new_json = json.dumps(sats.dict()) + if row.sats_pricing != new_json: + row.sats_pricing = new_json + s.add(row) + changed += 1 + except Exception as per_row_error: + logger.error( + "Failed to update pricing for model", + extra={ + "model_id": row.id, + "error": str(per_row_error), + "error_type": type(per_row_error).__name__, + }, + ) + if changed: + await s.commit() + + if updated_count > 0 or changed > 0: + logger.info( + "Updated sats pricing", + extra={ + "provider_models_updated": updated_count, + "database_overrides_updated": changed, + }, + ) + + +async def update_sats_pricing() -> None: + """Periodically update sats pricing for all provider models and database overrides.""" + try: + if not settings.enable_pricing_refresh: + return + except Exception: + pass + + await _update_sats_pricing_once() + + while True: try: interval = getattr(settings, "pricing_refresh_interval_seconds", 120) jitter = max(0.0, float(interval) * 0.1) @@ -399,6 +572,19 @@ async def update_sats_pricing() -> None: except asyncio.CancelledError: break + try: + try: + if not settings.enable_pricing_refresh: + return + except Exception: + pass + + await _update_sats_pricing_once() + except asyncio.CancelledError: + break + except Exception as e: + logger.error(f"Error updating sats pricing: {e}") + async def refresh_models_periodically() -> None: """Background task: periodically fetch OpenRouter models and insert new ones. @@ -473,5 +659,8 @@ async def refresh_models_periodically() -> None: @models_router.get("/v1/models") @models_router.get("/models", include_in_schema=False) async def models(session: AsyncSession = Depends(get_session)) -> dict: - items = await list_models(session) + """Get all available models from all providers with database overrides applied.""" + from ..proxy import get_unique_models + + items = get_unique_models() return {"data": items} diff --git a/routstr/proxy.py b/routstr/proxy.py index 6f8d976e..736d8bc4 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -1,33 +1,173 @@ import json +from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import Response, StreamingResponse +from sqlmodel import select from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key from .core import get_logger -from .core.db import ApiKey, AsyncSession, get_session -from .core.settings import settings +from .core.db import ApiKey, AsyncSession, ModelRow, create_session, get_session from .payment.helpers import ( calculate_discounted_max_cost, check_token_balance, create_error_response, get_max_cost_for_model, ) -from .upstream import init_upstreams +from .payment.models import Model, _row_to_model +from .upstream import UpstreamProvider, init_upstreams, resolve_model_alias logger = get_logger(__name__) proxy_router = APIRouter() -upstreams = init_upstreams(settings.upstream_base_url, settings.upstream_api_key) -upstream = upstreams[0] +_upstreams: list[UpstreamProvider] = [] +_model_instances: dict[str, Model] = {} # All aliases -> Model +_provider_map: dict[str, UpstreamProvider] = {} # All aliases -> Provider +_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates) + + +async def initialize_upstreams() -> None: + """Initialize upstream providers from database during application startup.""" + global _upstreams + _upstreams = await init_upstreams() + logger.info(f"Initialized {len(_upstreams)} upstream providers") + await refresh_model_maps() + + +async def reinitialize_upstreams() -> None: + """Re-initialize upstream providers from database (called after admin changes).""" + global _upstreams + _upstreams = await init_upstreams() + logger.info( + "Re-initialized upstream providers from admin action", + extra={"provider_count": len(_upstreams)}, + ) + await refresh_model_maps() + + +def get_upstreams() -> list[UpstreamProvider]: + """Get the initialized upstream providers. + + Returns: + List of upstream provider instances + """ + return _upstreams + + +def get_model_instance(model_id: str) -> Model | None: + """Get Model instance by ID from global cache.""" + return _model_instances.get(model_id) + + +def get_provider_for_model(model_id: str) -> UpstreamProvider | None: + """Get UpstreamProvider for model ID from global cache.""" + return _provider_map.get(model_id) + + +def get_unique_models() -> list[Model]: + """Get list of unique models (no duplicates from aliases).""" + return list(_unique_models.values()) + + +async def refresh_model_maps() -> None: + """Refresh global model and provider maps in-place.""" + global _model_instances, _provider_map, _unique_models + + model_instances: dict[str, Model] = {} + provider_map: dict[str, UpstreamProvider] = {} + unique_models: dict[str, Model] = {} + openrouter: UpstreamProvider | None = None + other_upstreams: list[UpstreamProvider] = [] + + async with create_session() as session: + result = await session.exec(select(ModelRow).where(ModelRow.enabled)) + override_rows = result.all() + overrides_by_id = { + row.id: row for row in override_rows if row.upstream_provider_id is not None + } + + for upstream in _upstreams: + if upstream.base_url == "https://openrouter.ai/api/v1": + openrouter = upstream + else: + other_upstreams.append(upstream) + + def get_base_model_id(model_id: str) -> str: + """Get base model ID by removing provider prefix.""" + return model_id.split("/", 1)[1] if "/" in model_id else model_id + + if openrouter: + for model in openrouter.get_cached_models(): + if model.enabled: + model_to_use = ( + _row_to_model(overrides_by_id[model.id]) + if model.id in overrides_by_id + else model + ) + base_id = get_base_model_id(model_to_use.id) + if base_id not in unique_models: + unique_models[base_id] = model_to_use + for alias in resolve_model_alias(model.id, model_to_use.canonical_slug): + model_instances[alias] = model_to_use + provider_map[alias] = openrouter + + for upstream in other_upstreams: + upstream_prefix = getattr(upstream, "upstream_name", None) + for model in upstream.get_cached_models(): + if model.enabled: + model_to_use = ( + _row_to_model(overrides_by_id[model.id]) + if model.id in overrides_by_id + else model + ) + base_id = get_base_model_id(model_to_use.id) + unique_models[base_id] = model_to_use + + aliases = resolve_model_alias(model.id, model_to_use.canonical_slug) + + if upstream_prefix and "/" not in model.id: + prefixed_id = f"{upstream_prefix}/{model.id}" + if prefixed_id not in aliases: + aliases.append(prefixed_id) + + for alias in aliases: + model_instances[alias] = model_to_use + provider_map[alias] = upstream + + _model_instances = model_instances + _provider_map = provider_map + _unique_models = unique_models + + logger.debug( + "Refreshed model maps", + extra={ + "unique_model_count": len(_unique_models), + "total_alias_count": len(_model_instances), + }, + ) + + +async def refresh_model_maps_periodically() -> None: + """Background task to refresh model maps every minute.""" + import asyncio + + while True: + try: + await asyncio.sleep(60) + await refresh_model_maps() + except asyncio.CancelledError: + break + except Exception as e: + logger.error( + "Error refreshing model maps", + extra={"error": str(e), "error_type": type(e).__name__}, + ) @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: - """Main proxy endpoint handler.""" - request_body = await request.body() headers = dict(request.headers) if "x-cashu" not in headers and "authorization" not in headers.keys(): @@ -35,7 +175,7 @@ async def proxy( "unauthorized", "Unauthorized", 401, request=request ) - logger.info( + logger.info( # TODO: move to middleware, async "Received proxy request", extra={ "method": request.method, @@ -45,76 +185,47 @@ async def proxy( }, ) - # Parse JSON body if present, handle empty/invalid JSON - request_body_dict = {} - if request_body: - try: - request_body_dict = json.loads(request_body) - logger.debug( - "Request body parsed", - extra={ - "path": path, - "body_keys": list(request_body_dict.keys()), - "model": request_body_dict.get("model", "not_specified"), - }, - ) - except json.JSONDecodeError as e: - logger.error( - "Invalid JSON in request body", - extra={ - "error": str(e), - "path": path, - "body_preview": request_body[:200].decode(errors="ignore") - if request_body - else "empty", - }, - ) - return Response( - content=json.dumps( - {"error": {"type": "invalid_request_error", "code": "invalid_json"}} - ), - status_code=400, - media_type="application/json", - ) + request_body = await request.body() + request_body_dict = parse_request_body_json(request_body, path) - model = request_body_dict.get("model", "unknown") - _max_cost_for_model = await get_max_cost_for_model(model=model, session=session) + model_id = request_body_dict.get("model", "unknown") + + model_obj = get_model_instance(model_id) + if not model_obj: + return create_error_response( + "invalid_model", f"Model '{model_id}' not found", 400, request=request + ) + + upstream = get_provider_for_model(model_id) + if not upstream: + return create_error_response( + "invalid_model", + f"No provider found for model '{model_id}'", + 400, + request=request, + ) + + _max_cost_for_model = await get_max_cost_for_model( + model=model_id, session=session, model_obj=model_obj + ) max_cost_for_model = await calculate_discounted_max_cost( _max_cost_for_model, request_body_dict, session ) check_token_balance(headers, request_body_dict, max_cost_for_model) - # Handle authentication if x_cashu := headers.get("x-cashu", None): - logger.info( - "Processing X-Cashu payment", - extra={ - "path": path, - "token_preview": x_cashu[:20] + "..." if len(x_cashu) > 20 else x_cashu, - }, - ) return await upstream.handle_x_cashu(request, x_cashu, path, max_cost_for_model) elif auth := headers.get("authorization", None): - logger.debug( - "Processing bearer token authentication", - extra={ - "path": path, - "token_preview": auth[:20] + "..." if len(auth) > 20 else auth, - }, - ) key = await get_bearer_token_key(headers, path, session, auth) else: if request.method not in ["GET"]: - logger.warning( - "Unauthorized request - no authentication provided", - extra={"method": request.method, "path": path}, - ) - return Response( - content=json.dumps({"detail": "Unauthorized"}), + raise HTTPException( status_code=401, - media_type="application/json", + detail={ + "error": {"type": "invalid_request_error", "code": "unauthorized"} + }, ) logger.debug("Processing unauthenticated GET request", extra={"path": path}) @@ -124,38 +235,13 @@ async def proxy( # Only pay for request if we have request body data (for completions endpoints) if request_body_dict: - logger.info( - "Processing payment for request", - extra={ - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "key_balance_before": key.balance, - "model": request_body_dict.get("model", "unknown"), - }, - ) - try: await pay_for_request(key, max_cost_for_model, session) - logger.info( - "Payment processed successfully", - extra={ - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "key_balance_after": key.balance, - "model": request_body_dict.get("model", "unknown"), - }, + except Exception: + raise HTTPException( + status_code=402, + detail={"error": {"type": "payment_error", "code": "payment_error"}}, ) - except Exception as e: - logger.error( - "Payment processing failed", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "path": path, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - raise # Prepare headers for upstream headers = upstream.prepare_headers(dict(request.headers)) @@ -270,3 +356,37 @@ async def get_bearer_token_key( }, ) raise + + +def parse_request_body_json(request_body: bytes, path: str) -> dict[str, Any]: + request_body_dict = {} + if request_body: + try: + request_body_dict = json.loads(request_body) + logger.debug( + "Request body parsed", + extra={ + "path": path, + "body_keys": list(request_body_dict.keys()), + "model": request_body_dict.get("model", "not_specified"), + }, + ) + except json.JSONDecodeError as e: + logger.error( + "Invalid JSON in request body", + extra={ + "error": str(e), + "path": path, + "body_preview": request_body[:200].decode(errors="ignore") + if request_body + else "empty", + }, + ) + raise HTTPException( + status_code=400, + detail={ + "error": {"type": "invalid_request_error", "code": "invalid_json"} + }, + ) + + return request_body_dict diff --git a/routstr/upstream.py b/routstr/upstream.py index d9728017..d9edfa1a 100644 --- a/routstr/upstream.py +++ b/routstr/upstream.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Mapping import httpx if TYPE_CHECKING: + from .core.settings import Settings from .payment.cost_caculation import CostData, MaxCostData from fastapi import BackgroundTasks, HTTPException, Request @@ -16,50 +17,415 @@ 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.db import ApiKey, AsyncSession, ModelRow, UpstreamProviderRow, create_session from .payment.helpers import create_error_response -from .payment.models import Model +from .payment.models import Model, async_fetch_openrouter_models logger = get_logger(__name__) -def init_upstreams( - base_url: str, api_key: str, api_version: str | None = None -) -> list[UpstreamProvider]: - """Initialize upstream providers based on settings. +def resolve_model_alias(model_id: str, canonical_slug: str | None = None) -> list[str]: + """Resolve model ID to all possible aliases. + + Returns list of aliases including canonical slug and variations without provider prefix. Args: - base_url: Base URL of the upstream API endpoint - api_key: API key for authenticating with the upstream service - api_version: API version for Azure OpenAI + model_id: Model identifier (e.g., "gpt-5-mini" or "openai/gpt-5-mini") + canonical_slug: Optional canonical slug from provider (e.g., "openai/gpt-5-pro-2025-10-06") + + Returns: + List of possible model ID aliases """ + aliases = [model_id] + + base_model = model_id + if "/" in model_id: + without_prefix = model_id.split("/", 1)[1] + aliases.append(without_prefix) + base_model = without_prefix + + date_pattern = re.compile(r"-\d{4}-\d{2}-\d{2}$") + if date_pattern.search(base_model): + base_without_date = date_pattern.sub("", base_model) + if base_without_date not in aliases: + aliases.append(base_without_date) + if "/" in model_id: + prefix = model_id.split("/", 1)[0] + prefixed_without_date = f"{prefix}/{base_without_date}" + if prefixed_without_date not in aliases: + aliases.append(prefixed_without_date) + + if canonical_slug and canonical_slug not in aliases: + aliases.append(canonical_slug) + if "/" in canonical_slug: + canonical_without_prefix = canonical_slug.split("/", 1)[1] + if canonical_without_prefix not in aliases: + aliases.append(canonical_without_prefix) + if date_pattern.search(canonical_without_prefix): + canonical_base = date_pattern.sub("", canonical_without_prefix) + if canonical_base not in aliases: + aliases.append(canonical_base) + + return aliases + + +async def get_all_models_with_overrides( + upstreams: list[UpstreamProvider], +) -> list[Model]: + """Get all models from all providers with database overrides applied. + + Models in the database with upstream_provider_id set are treated as overrides + that replace the provider's model with the same ID. + + Args: + upstreams: List of upstream provider instances + + Returns: + List of Model objects with overrides applied + """ + from sqlmodel import select + + from .payment.models import _row_to_model + + async with create_session() as session: + result = await session.exec(select(ModelRow).where(ModelRow.enabled)) + override_rows = result.all() + overrides_by_id = { + row.id: row for row in override_rows if row.upstream_provider_id is not None + } + + all_models: dict[str, Model] = {} + + for upstream in upstreams: + for model in upstream.get_cached_models(): + if model.id in overrides_by_id: + all_models[model.id] = _row_to_model(overrides_by_id[model.id]) + elif model.enabled: + all_models[model.id] = model + + return list(all_models.values()) + + +async def get_model_with_override( + model_id: str, + upstreams: list[UpstreamProvider], +) -> Model | None: + """Get a specific model from providers with database override applied. + + Resolves model aliases automatically (e.g., both "gpt-5-mini" and "openai/gpt-5-mini"). + + Args: + model_id: Model identifier (with or without provider prefix) + upstreams: List of upstream provider instances + + Returns: + Model object or None if not found + """ + from sqlmodel import select + + from .payment.models import _row_to_model + + aliases = resolve_model_alias(model_id) + + async with create_session() as session: + for alias in aliases: + result = await session.exec( + select(ModelRow).where( + ModelRow.id == alias, + ModelRow.upstream_provider_id.isnot(None), # type: ignore + ModelRow.enabled, + ) + ) + override_row = result.first() + if override_row: + return _row_to_model(override_row) + + for alias in aliases: + for upstream in upstreams: + model = upstream.get_cached_model_by_id(alias) + if model and model.enabled: + return model + + return None + + +async def refresh_upstreams_models_periodically( + upstreams: list[UpstreamProvider], +) -> None: + """Background task to periodically refresh models cache for all providers. + + Args: + upstreams: List of upstream provider instances + """ + import asyncio + import random + from .core.settings import settings - upstreams: list[UpstreamProvider] = [] - if settings.chat_completions_api_version: - upstreams.append( - AzureUpstreamProvider( - settings.upstream_base_url, - settings.upstream_api_key, - settings.chat_completions_api_version, + interval = getattr(settings, "models_refresh_interval_seconds", 0) + if not interval or interval <= 0: + logger.info("Provider models refresh disabled (interval <= 0)") + return + + while True: + try: + for upstream in upstreams: + try: + await upstream.refresh_models_cache() + except Exception as e: + logger.error( + f"Error refreshing models for {upstream.upstream_name or upstream.base_url}", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + except asyncio.CancelledError: + break + except Exception as e: + logger.error( + "Error in provider models refresh loop", + extra={"error": str(e), "error_type": type(e).__name__}, ) + + try: + jitter = max(0.0, float(interval) * 0.1) + await asyncio.sleep(interval + random.uniform(0, jitter)) + except asyncio.CancelledError: + break + + +import os + + +async def init_upstreams() -> list[UpstreamProvider]: + """Initialize upstream providers from database. + + Seeds database with providers from settings if empty, then loads and instantiates + provider instances from database records, and refreshes their models cache. + """ + from sqlmodel import select + + from .core.settings import settings + + async with create_session() as session: + result = await session.exec(select(UpstreamProviderRow)) + existing_providers = result.all() + + if not existing_providers: + logger.info( + "No upstream providers found in database, seeding from settings" + ) + await _seed_providers_from_settings(session, settings) + await session.commit() + result = await session.exec(select(UpstreamProviderRow)) + existing_providers = result.all() + + upstreams: list[UpstreamProvider] = [] + for provider_row in existing_providers: + if not provider_row.enabled: + logger.debug(f"Skipping disabled provider: {provider_row.base_url}") + continue + + provider = _instantiate_provider(provider_row) + if provider: + await provider.refresh_models_cache() + upstreams.append(provider) + logger.info( + f"Initialized {provider_row.provider_type} provider", + extra={ + "base_url": provider_row.base_url, + "models_cached": len(provider.get_cached_models()), + }, + ) + + return upstreams + + +async def _seed_providers_from_settings( + session: AsyncSession, settings: "Settings" +) -> None: + """Seed database with upstream providers from environment variables. + + Args: + session: Database session + """ + from sqlmodel import select + + from .core.settings import settings + + providers_to_add: list[UpstreamProviderRow] = [] + seeded_base_urls: set[str] = set() + + openai_api_key = os.environ.get("OPENAI_API_KEY") + if openai_api_key: + base_url = "https://api.openai.com/v1" + result = await session.exec( + select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url) + ) + if not result.first(): + providers_to_add.append( + UpstreamProviderRow( + provider_type="openai", + base_url=base_url, + api_key=openai_api_key, + enabled=True, + ) + ) + seeded_base_urls.add(base_url) + + anthropic_api_key = os.environ.get("ANTHROPIC_API_KEY") + if anthropic_api_key: + base_url = "https://api.anthropic.com/v1" + result = await session.exec( + select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url) + ) + if not result.first(): + providers_to_add.append( + UpstreamProviderRow( + provider_type="anthropic", + base_url=base_url, + api_key=anthropic_api_key, + enabled=True, + ) + ) + seeded_base_urls.add(base_url) + + openrouter_api_key = os.environ.get("OPENROUTER_API_KEY") + if openrouter_api_key: + base_url = "https://openrouter.ai/api/v1" + result = await session.exec( + select(UpstreamProviderRow).where(UpstreamProviderRow.base_url == base_url) + ) + if not result.first(): + providers_to_add.append( + UpstreamProviderRow( + provider_type="openrouter", + base_url=base_url, + api_key=openrouter_api_key, + enabled=True, + ) + ) + seeded_base_urls.add(base_url) + + if settings.chat_completions_api_version and settings.upstream_base_url: + base_url = settings.upstream_base_url + if base_url not in seeded_base_urls: + result = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.base_url == base_url + ) + ) + if not result.first(): + providers_to_add.append( + UpstreamProviderRow( + provider_type="azure", + base_url=base_url, + api_key=settings.upstream_api_key, + api_version=settings.chat_completions_api_version, + enabled=True, + ) + ) + seeded_base_urls.add(base_url) + + if settings.upstream_base_url and settings.upstream_api_key: + base_url = settings.upstream_base_url + if base_url not in seeded_base_urls: + result = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.base_url == base_url + ) + ) + if not result.first(): + if "api.openai.com" in base_url.lower(): + providers_to_add.append( + UpstreamProviderRow( + provider_type="openai", + base_url=base_url, + api_key=settings.upstream_api_key, + enabled=True, + ) + ) + elif "openrouter.ai/api/v1" in base_url.lower(): + providers_to_add.append( + UpstreamProviderRow( + provider_type="openrouter", + base_url=base_url, + api_key=settings.upstream_api_key, + enabled=True, + ) + ) + else: + providers_to_add.append( + UpstreamProviderRow( + provider_type="generic", + base_url=base_url, + api_key=settings.upstream_api_key, + enabled=True, + ) + ) + seeded_base_urls.add(base_url) + + for provider in providers_to_add: + session.add(provider) + logger.info( + f"Seeding {provider.provider_type} provider", + extra={"base_url": provider.base_url}, ) - if "api.openai.com" in settings.upstream_base_url.lower(): - upstreams.append(OpenAIUpstreamProvider(settings.upstream_api_key)) - elif "openrouter.ai/api/v1" in settings.upstream_base_url.lower(): - upstreams.append(OpenRouterUpstreamProvider(settings.upstream_api_key)) - else: - upstreams.append( - UpstreamProvider(settings.upstream_base_url, settings.upstream_api_key) - ) - return upstreams +def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider | None: + """Instantiate an UpstreamProvider from a database row. + + Args: + provider_row: Database row containing provider configuration + + Returns: + Instantiated provider or None if provider type is unknown + """ + try: + if provider_row.provider_type == "openai": + return OpenAIUpstreamProvider(provider_row.api_key) + elif provider_row.provider_type == "azure": + if not provider_row.api_version: + logger.error( + "Azure provider missing api_version", + extra={"base_url": provider_row.base_url}, + ) + return None + return AzureUpstreamProvider( + provider_row.base_url, + provider_row.api_key, + provider_row.api_version, + ) + elif provider_row.provider_type == "openrouter": + return OpenRouterUpstreamProvider(provider_row.api_key) + elif provider_row.provider_type == "generic": + return UpstreamProvider(provider_row.base_url, provider_row.api_key) + else: + logger.error( + f"Unknown provider type: {provider_row.provider_type}", + extra={"base_url": provider_row.base_url}, + ) + return None + except Exception as e: + logger.error( + f"Failed to instantiate provider: {e}", + extra={ + "provider_type": provider_row.provider_type, + "base_url": provider_row.base_url, + "error": str(e), + }, + ) + return None class UpstreamProvider: """Provider for forwarding requests to an upstream AI service API.""" + base_url: str + api_key: str + upstream_name: str | None = None + _models_cache: list[Model] = [] + _models_by_id: dict[str, Model] = {} + def __init__(self, base_url: str, api_key: str): """Initialize the upstream provider. @@ -69,6 +435,8 @@ class UpstreamProvider: """ self.base_url = base_url self.api_key = api_key + self._models_cache = [] + self._models_by_id = {} def prepare_headers(self, request_headers: dict) -> dict: """Prepare headers for upstream request by removing proxy-specific headers and adding authentication. @@ -136,6 +504,60 @@ class UpstreamProvider: """ return query_params or {} + def transform_model_name(self, model_id: str) -> str: + """Transform model ID for this provider's API format. + + Base implementation returns model_id unchanged. Override in subclasses for provider-specific transformations. + + Args: + model_id: Model identifier (may include provider prefix) + + Returns: + Transformed model ID for this provider + """ + return model_id + + def prepare_request_body(self, body: bytes | None) -> bytes | None: + """Transform request body for provider-specific requirements. + + Automatically transforms model names in the request body. + + Args: + body: Original request body bytes + + Returns: + Transformed request body bytes + """ + if not body: + return body + + try: + data = json.loads(body) + if isinstance(data, dict) and "model" in data: + original_model = data["model"] + transformed_model = self.transform_model_name(original_model) + if transformed_model != original_model: + data["model"] = transformed_model + logger.debug( + "Transformed model name in request", + extra={ + "original": original_model, + "transformed": transformed_model, + "provider": self.upstream_name or self.base_url, + }, + ) + return json.dumps(data).encode() + except Exception as e: + logger.debug( + "Could not transform request body", + extra={ + "error": str(e), + "provider": self.upstream_name or self.base_url, + }, + ) + + return body + def _extract_upstream_error_message( self, body_bytes: bytes ) -> tuple[str, str | None]: @@ -560,6 +982,8 @@ class UpstreamProvider: url = f"{self.base_url}/{path}" + transformed_body = self.prepare_request_body(request_body) + logger.info( "Forwarding request to upstream", extra={ @@ -578,13 +1002,13 @@ class UpstreamProvider: ) try: - if request_body is not None: + if transformed_body is not None: response = await client.send( client.build_request( request.method, url, headers=headers, - content=request_body, + content=transformed_body, params=self.prepare_params(path, request.query_params), ), stream=True, @@ -1326,6 +1750,9 @@ class UpstreamProvider: url = f"{self.base_url}/{path}" + request_body = await request.body() + transformed_body = self.prepare_request_body(request_body) + logger.debug( "Forwarding request to upstream", extra={ @@ -1347,7 +1774,7 @@ class UpstreamProvider: request.method, url, headers=headers, - content=request.stream(), + content=transformed_body if transformed_body else request_body, params=self.prepare_params(path, request.query_params), ), stream=True, @@ -1546,12 +1973,66 @@ class UpstreamProvider: token=x_cashu_token, ) + async def fetch_models(self) -> list[Model]: + """Fetch available models from upstream API and update cache. + + Returns: + List of Model objects with pricing + """ + logger.debug(f"Fetching models for {self.upstream_name or self.base_url}") + return [] + + async def refresh_models_cache(self) -> None: + """Refresh the in-memory models cache from upstream API.""" + try: + models = await self.fetch_models() + self._models_cache = models + self._models_by_id = {m.id: m for m in models} + logger.info( + f"Refreshed models cache for {self.upstream_name or self.base_url}", + extra={"model_count": len(models)}, + ) + except Exception as e: + logger.error( + f"Failed to refresh models cache for {self.upstream_name or self.base_url}", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + + def get_cached_models(self) -> list[Model]: + """Get cached models for this provider. + + Returns: + List of cached Model objects + """ + return self._models_cache + + def get_cached_model_by_id(self, model_id: str) -> Model | None: + """Get a specific cached model by ID. + + Args: + model_id: Model identifier + + Returns: + Model object or None if not found + """ + return self._models_by_id.get(model_id) + class OpenAIUpstreamProvider(UpstreamProvider): """Upstream provider specifically configured for OpenAI API.""" def __init__(self, api_key: str): - super().__init__(base_url="https://api.openai.com", api_key=api_key) + self.upstream_name = "openai" + super().__init__(base_url="https://api.openai.com/v1", api_key=api_key) + + def transform_model_name(self, model_id: str) -> str: + """Strip 'openai/' prefix for OpenAI API compatibility.""" + return model_id.removeprefix("openai/") + + async def fetch_models(self) -> list[Model]: + """Fetch OpenAI models from OpenRouter API filtered by openai source.""" + models_data = await async_fetch_openrouter_models(source_filter="openai") + return [Model(**model) for model in models_data] # type: ignore class AzureUpstreamProvider(UpstreamProvider): @@ -1595,32 +2076,10 @@ class OpenRouterUpstreamProvider(UpstreamProvider): Args: api_key: OpenRouter API key for authentication """ + self.upstream_name = "openrouter" super().__init__(base_url="https://openrouter.ai/api/v1", api_key=api_key) - async def fetch_models(self) -> dict: - """Fetch available models from OpenRouter API. - - Returns: - Raw JSON response containing model data - """ - async with httpx.AsyncClient() as client: - response = await client.get( - "https://openrouter.ai/api/v1/models", - headers={"Authorization": f"Bearer {self.api_key}"}, - ) - return response.json() - - async def models(self) -> list[Model]: - """Get list of available models from OpenRouter. - - Returns: - List of Model objects representing available models - """ - response_data = await self.fetch_models() - models_list: list[Model] = [] - for model_data in response_data.get("data", []): - try: - models_list.append(Model(**model_data)) # type: ignore - except Exception: - continue - return models_list + async def fetch_models(self) -> list[Model]: + """Fetch all OpenRouter models.""" + models_data = await async_fetch_openrouter_models() + return [Model(**model) for model in models_data] # type: ignore From 0da08fb94593ecf7910b6bf23f6e7d6bf8200ab8 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Mon, 20 Oct 2025 12:45:02 +0800 Subject: [PATCH 10/95] refactor pricing, provider fees, realtime model map updates --- .../a1a1a1a1a1a1_composite_pk_for_models.py | 64 +++ ...b9c0d1e2_add_fees_to_upstream_providers.py | 27 + routstr/core/admin.py | 498 +++++++++++++----- routstr/core/db.py | 16 +- routstr/core/main.py | 12 +- routstr/payment/cost_caculation.py | 7 +- routstr/payment/helpers.py | 14 +- routstr/payment/models.py | 304 +++++------ routstr/payment/price.py | 95 ++-- routstr/proxy.py | 66 ++- routstr/upstream.py | 175 ++++-- 11 files changed, 854 insertions(+), 424 deletions(-) create mode 100644 migrations/versions/a1a1a1a1a1a1_composite_pk_for_models.py create mode 100644 migrations/versions/f7a8b9c0d1e2_add_fees_to_upstream_providers.py diff --git a/migrations/versions/a1a1a1a1a1a1_composite_pk_for_models.py b/migrations/versions/a1a1a1a1a1a1_composite_pk_for_models.py new file mode 100644 index 00000000..0cbafc26 --- /dev/null +++ b/migrations/versions/a1a1a1a1a1a1_composite_pk_for_models.py @@ -0,0 +1,64 @@ +"""change models to composite primary key (id, upstream_provider_id) + +Revision ID: a1a1a1a1a1a1 +Revises: f7a8b9c0d1e2 +Create Date: 2025-10-20 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "a1a1a1a1a1a1" +down_revision = "f7a8b9c0d1e2" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + + if "models" in inspector.get_table_names(): + op.drop_table("models") + + op.create_table( + "models", + sa.Column("id", sa.String(), nullable=False), + sa.Column("upstream_provider_id", sa.Integer(), nullable=False), + sa.Column("name", sa.String(), nullable=False), + sa.Column("created", sa.Integer(), nullable=False), + sa.Column("description", sa.Text(), nullable=False), + sa.Column("context_length", sa.Integer(), nullable=False), + sa.Column("architecture", sa.Text(), nullable=False), + sa.Column("pricing", sa.Text(), nullable=False), + sa.Column("sats_pricing", sa.Text(), nullable=True), + sa.Column("per_request_limits", sa.Text(), nullable=True), + sa.Column("top_provider", sa.Text(), nullable=True), + sa.Column("enabled", sa.Boolean(), nullable=False, server_default="1"), + sa.PrimaryKeyConstraint("id", "upstream_provider_id"), + sa.ForeignKeyConstraint( + ["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE" + ), + ) + + +def downgrade() -> None: + op.drop_table("models") + op.create_table( + "models", + sa.Column("id", sa.String(), primary_key=True, nullable=False), + sa.Column("name", sa.String(), nullable=False), + sa.Column("created", sa.Integer(), nullable=False), + sa.Column("description", sa.Text(), nullable=False), + sa.Column("context_length", sa.Integer(), nullable=False), + sa.Column("architecture", sa.Text(), nullable=False), + sa.Column("pricing", sa.Text(), nullable=False), + sa.Column("sats_pricing", sa.Text(), nullable=True), + sa.Column("per_request_limits", sa.Text(), nullable=True), + sa.Column("top_provider", sa.Text(), nullable=True), + sa.Column("enabled", sa.Boolean(), nullable=False, server_default="1"), + sa.Column("upstream_provider_id", sa.Integer(), nullable=True), + sa.ForeignKeyConstraint(["upstream_provider_id"], ["upstream_providers.id"]), + ) diff --git a/migrations/versions/f7a8b9c0d1e2_add_fees_to_upstream_providers.py b/migrations/versions/f7a8b9c0d1e2_add_fees_to_upstream_providers.py new file mode 100644 index 00000000..2c921094 --- /dev/null +++ b/migrations/versions/f7a8b9c0d1e2_add_fees_to_upstream_providers.py @@ -0,0 +1,27 @@ +"""add provider_fee to upstream_providers + +Revision ID: f7a8b9c0d1e2 +Revises: e1f2a3b4c5d6 +Create Date: 2025-10-13 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "f7a8b9c0d1e2" +down_revision = "e1f2a3b4c5d6" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column( + "upstream_providers", + sa.Column("provider_fee", sa.Float(), nullable=False, server_default="1.01"), + ) + + +def downgrade() -> None: + op.drop_column("upstream_providers", "provider_fee") diff --git a/routstr/core/admin.py b/routstr/core/admin.py index bb54ffe7..288b1c8e 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1389,30 +1389,30 @@ UPSTREAM_PROVIDERS_JS: str = """ """ @@ -1936,7 +2083,6 @@ def upstream_providers_page() -> str: - @@ -1944,7 +2090,7 @@ def upstream_providers_page() -> str: - +
ID Type Base URL Status
Loading…
Loading…
@@ -1958,7 +2104,15 @@ def upstream_providers_page() -> str:
Loading models…
@@ -2003,7 +2159,7 @@ def upstream_providers_page() -> str: - + @@ -2017,6 +2173,10 @@ def upstream_providers_page() -> str: + + + Leave empty to use default (OpenRouter: 1.06, Others: 1.01) +