From 71623b121f729ba83ee2d2b5ee6b76d5ee45e561 Mon Sep 17 00:00:00 2001 From: 9qeklajc <9qeklajc> Date: Sun, 29 Jun 2025 12:17:41 +0200 Subject: [PATCH] add simple forward func --- router/proxy.py | 55 ++++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 50 insertions(+), 5 deletions(-) diff --git a/router/proxy.py b/router/proxy.py index 50292d81..5de67a94 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -1,6 +1,7 @@ import json import os import re +import traceback from typing import AsyncGenerator import httpx @@ -168,7 +169,7 @@ async def forward_to_upstream( path: str, headers: dict, request_body: bytes | None, - key: ApiKey | None, + key: ApiKey, session: AsyncSession, ) -> Response | StreamingResponse: """Forward request to upstream and handle the response.""" @@ -302,6 +303,9 @@ async def proxy( media_type="application/json", ) + # Prepare headers for upstream + headers = prepare_upstream_headers(dict(request.headers)) + # Handle authentication if x_cashu := headers.get("x-cashu", None): # Check token balance before authentication for cashu tokens @@ -313,21 +317,18 @@ async def proxy( key = await get_bearer_token_key(headers, path, session, auth) else: - key = None if request.method not in ["GET"]: return Response( content=json.dumps({"detail": "Unauthorized"}), status_code=401, media_type="application/json", ) + return await forward_get_to_upstream(request, path, headers) # Only pay for request if we have request body data (for completions endpoints) if request_body_dict and key is not None: await pay_for_request(key, session, request_body_dict) - # Prepare headers for upstream - headers = prepare_upstream_headers(dict(request.headers)) - # Forward to upstream and handle response response = await forward_to_upstream( request, path, headers, request_body, key, session @@ -392,3 +393,47 @@ async def get_bearer_token_key( refund_address, key_expiry_time, # type: ignore ) + + +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"{UPSTREAM_BASE_URL}/{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=request.query_params, + ), + ) + + return StreamingResponse( + response.aiter_bytes(), + status_code=response.status_code, + headers=dict(response.headers), + ) + except Exception as exc: + tb = traceback.format_exc() + print( + f"Unexpected error: {exc}\n" + f"Request details: method={request.method}, url={url}, headers={headers}, " + f"path={path}, query_params={dict(request.query_params)}\n" + f"Traceback:\n{tb}" + ) + return create_error_response( + "internal_error", "An unexpected server error occurred", 500 + )