add simple forward func

This commit is contained in:
9qeklajc
2025-06-29 12:17:41 +02:00
parent 708b79d979
commit 71623b121f
+50 -5
View File
@@ -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
)