mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-11 19:57:32 +00:00
add simple forward func
This commit is contained in:
+50
-5
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user