mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
x_cashu mvp + proxy refactor
This commit is contained in:
@@ -203,6 +203,14 @@ async def refund_balance(amount_msats: int, key: ApiKey, session: AsyncSession)
|
||||
)
|
||||
|
||||
|
||||
async def x_cashu_refund(key: ApiKey, session: AsyncSession) -> str:
|
||||
async with WALLET_LOCK:
|
||||
refund_token = await WALLET.send(key.balance)
|
||||
await session.delete(key)
|
||||
await session.commit()
|
||||
return refund_token
|
||||
|
||||
|
||||
async def redeem(cashu_token: str, lnurl: str) -> int:
|
||||
async with WALLET_LOCK:
|
||||
amount_sats, _ = await WALLET.redeem(cashu_token)
|
||||
|
||||
+264
-186
@@ -1,15 +1,16 @@
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from hashlib import sha256
|
||||
from typing import AsyncGenerator
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, Request
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
|
||||
from .auth import adjust_payment_for_tokens, pay_for_request, validate_bearer_key
|
||||
from .cashu import pay_out
|
||||
from .db import AsyncSession, create_session, get_session
|
||||
from .cashu import x_cashu_refund
|
||||
from .db import ApiKey, AsyncSession, create_session, get_session
|
||||
|
||||
UPSTREAM_BASE_URL = os.environ["UPSTREAM_BASE_URL"]
|
||||
UPSTREAM_API_KEY = os.environ.get("UPSTREAM_API_KEY", "")
|
||||
@@ -17,85 +18,68 @@ UPSTREAM_API_KEY = os.environ.get("UPSTREAM_API_KEY", "")
|
||||
proxy_router = APIRouter()
|
||||
|
||||
|
||||
@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:
|
||||
auth = request.headers.get("Authorization", "")
|
||||
bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else ""
|
||||
refund_address = request.headers.get("Refund-LNURL", None)
|
||||
key_expiry_time = request.headers.get("Key-Expiry-Time", None)
|
||||
async def validate_and_read_json_body(
|
||||
request: Request, path: str
|
||||
) -> tuple[bytes | None, Response | None]:
|
||||
"""Validate and read JSON body for requests that require it.
|
||||
|
||||
# Validate key_expiry_time header
|
||||
if key_expiry_time:
|
||||
try:
|
||||
key_expiry_time = int(key_expiry_time) # type: ignore
|
||||
except ValueError:
|
||||
return Response(
|
||||
content="Invalid Key-Expiry-Time: must be a valid Unix timestamp",
|
||||
status_code=400,
|
||||
)
|
||||
if not refund_address:
|
||||
return Response(
|
||||
content="Error: Refund-LNURL header required when using Key-Expiry-Time",
|
||||
status_code=400,
|
||||
)
|
||||
else:
|
||||
key_expiry_time = None
|
||||
Returns:
|
||||
tuple of (request_body, error_response)
|
||||
- If successful: (body_bytes, None)
|
||||
- If error: (None, error_response)
|
||||
"""
|
||||
if request.method not in ["POST", "PUT", "PATCH"] or not path.endswith(
|
||||
"chat/completions"
|
||||
):
|
||||
return None, None
|
||||
|
||||
key = await validate_bearer_key(
|
||||
bearer_key,
|
||||
session,
|
||||
refund_address,
|
||||
key_expiry_time, # type: ignore
|
||||
)
|
||||
|
||||
# Pre-validate JSON for requests that require it
|
||||
request_body = None
|
||||
if request.method in ["POST", "PUT", "PATCH"] and path.endswith("chat/completions"):
|
||||
try:
|
||||
request_body = await request.body()
|
||||
# Try to parse JSON to validate it
|
||||
if request_body:
|
||||
json.loads(request_body)
|
||||
except json.JSONDecodeError as e:
|
||||
return Response(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": f"Invalid JSON in request body: {str(e)}",
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_json",
|
||||
}
|
||||
try:
|
||||
request_body = await request.body()
|
||||
# Try to parse JSON to validate it
|
||||
if request_body:
|
||||
json.loads(request_body)
|
||||
return request_body, None
|
||||
except json.JSONDecodeError as e:
|
||||
return None, Response(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": f"Invalid JSON in request body: {str(e)}",
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_json",
|
||||
}
|
||||
),
|
||||
status_code=400,
|
||||
media_type="application/json",
|
||||
)
|
||||
except Exception:
|
||||
return Response(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": "Error reading request body",
|
||||
"type": "invalid_request_error",
|
||||
"code": "request_error",
|
||||
}
|
||||
}
|
||||
),
|
||||
status_code=400,
|
||||
media_type="application/json",
|
||||
)
|
||||
except Exception:
|
||||
return None, Response(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": "Error reading request body",
|
||||
"type": "invalid_request_error",
|
||||
"code": "request_error",
|
||||
}
|
||||
),
|
||||
status_code=400,
|
||||
media_type="application/json",
|
||||
)
|
||||
}
|
||||
),
|
||||
status_code=400,
|
||||
media_type="application/json",
|
||||
)
|
||||
|
||||
await pay_for_request(key, session, request, request_body)
|
||||
|
||||
# Prepare headers, removing sensitive/problematic ones
|
||||
headers = dict(request.headers)
|
||||
def prepare_upstream_headers(request_headers: dict) -> dict:
|
||||
"""Prepare headers for upstream request, removing sensitive/problematic ones."""
|
||||
headers = dict(request_headers)
|
||||
# Remove headers that shouldn't be forwarded
|
||||
headers.pop("host", None)
|
||||
headers.pop("content-length", None)
|
||||
headers.pop("refund-lnurl", None)
|
||||
headers.pop("key-expiry-time", None)
|
||||
headers.pop("x-cashu", None)
|
||||
|
||||
# Handle authorization
|
||||
if UPSTREAM_API_KEY:
|
||||
headers["Authorization"] = f"Bearer {UPSTREAM_API_KEY}"
|
||||
headers.pop("authorization", None)
|
||||
@@ -103,6 +87,134 @@ async def proxy(
|
||||
headers.pop("Authorization", None)
|
||||
headers.pop("authorization", None)
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
def create_error_response(error_type: str, message: str, status_code: int) -> Response:
|
||||
"""Create a standardized error response."""
|
||||
return Response(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": message,
|
||||
"type": error_type,
|
||||
"code": status_code,
|
||||
}
|
||||
}
|
||||
),
|
||||
status_code=status_code,
|
||||
media_type="application/json",
|
||||
)
|
||||
|
||||
|
||||
async def handle_streaming_chat_completion(
|
||||
response: httpx.Response, key: ApiKey, session: AsyncSession
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming chat completion responses with token-based pricing."""
|
||||
|
||||
async def stream_with_cost() -> AsyncGenerator[bytes, None]:
|
||||
# Store all chunks to analyze
|
||||
stored_chunks = []
|
||||
|
||||
async for chunk in response.aiter_bytes():
|
||||
# Store chunk for later analysis
|
||||
stored_chunks.append(chunk)
|
||||
|
||||
# Pass through each chunk to client
|
||||
yield chunk
|
||||
|
||||
# Process stored chunks to find usage data
|
||||
# Start from the end and work backwards
|
||||
for i in range(len(stored_chunks) - 1, -1, -1):
|
||||
chunk = stored_chunks[i]
|
||||
if not chunk or chunk == b"":
|
||||
continue
|
||||
|
||||
try:
|
||||
# Split by "data: " to get individual SSE events
|
||||
events = re.split(b"data: ", chunk)
|
||||
for event_data in events:
|
||||
if (
|
||||
not event_data
|
||||
or event_data.strip() == b"[DONE]"
|
||||
or event_data.strip() == b""
|
||||
):
|
||||
continue
|
||||
|
||||
try:
|
||||
data = json.loads(event_data)
|
||||
if (
|
||||
"usage" in data
|
||||
and data["usage"] is not None
|
||||
and isinstance(data["usage"], dict)
|
||||
):
|
||||
# Found usage data, calculate cost
|
||||
# Create a new session for this operation
|
||||
async with create_session() as new_session:
|
||||
# Re-fetch the key in the new session
|
||||
fresh_key = await new_session.get(
|
||||
key.__class__, key.hashed_key
|
||||
)
|
||||
if fresh_key:
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
fresh_key, data, new_session
|
||||
)
|
||||
# Format as SSE and yield
|
||||
cost_json = json.dumps({"cost": cost_data})
|
||||
yield f"data: {cost_json}\n\n".encode()
|
||||
break
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error processing streaming response for cost: {e}")
|
||||
|
||||
return StreamingResponse(
|
||||
stream_with_cost(),
|
||||
status_code=response.status_code,
|
||||
headers=dict(response.headers),
|
||||
)
|
||||
|
||||
|
||||
async def handle_non_streaming_chat_completion(
|
||||
response: httpx.Response, key: ApiKey, session: AsyncSession
|
||||
) -> Response:
|
||||
"""Handle non-streaming chat completion responses with token-based pricing."""
|
||||
try:
|
||||
content = await response.aread()
|
||||
response_json = json.loads(content)
|
||||
cost_data = await adjust_payment_for_tokens(key, response_json, session)
|
||||
response_json["cost"] = cost_data
|
||||
|
||||
response_headers = dict(response.headers)
|
||||
|
||||
# Remove Transfer-Encoding header to avoid conflict with Content-Length header in common nginx setups
|
||||
if "transfer-encoding" in response_headers:
|
||||
del response_headers["transfer-encoding"]
|
||||
|
||||
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:
|
||||
print(f"Failed to parse JSON from upstream response: {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
print(f"Error adjusting payment for tokens: {e}")
|
||||
raise
|
||||
|
||||
|
||||
async def forward_to_upstream(
|
||||
request: Request,
|
||||
path: str,
|
||||
headers: dict,
|
||||
request_body: bytes | None,
|
||||
key: ApiKey,
|
||||
session: AsyncSession,
|
||||
) -> Response | StreamingResponse:
|
||||
"""Forward request to upstream and handle the response."""
|
||||
if path.startswith("v1/"):
|
||||
path = path.replace("v1/", "")
|
||||
|
||||
@@ -145,103 +257,19 @@ async def proxy(
|
||||
|
||||
if is_streaming and response.status_code == 200:
|
||||
# Process streaming response and extract cost from the last chunk
|
||||
async def stream_with_cost() -> AsyncGenerator[bytes, None]:
|
||||
# Store all chunks to analyze
|
||||
stored_chunks = []
|
||||
|
||||
async for chunk in response.aiter_bytes():
|
||||
# Store chunk for later analysis
|
||||
stored_chunks.append(chunk)
|
||||
|
||||
# Pass through each chunk to client
|
||||
yield chunk
|
||||
|
||||
# Process stored chunks to find usage data
|
||||
# Start from the end and work backwards
|
||||
for i in range(len(stored_chunks) - 1, -1, -1):
|
||||
chunk = stored_chunks[i]
|
||||
if not chunk or chunk == b"":
|
||||
continue
|
||||
|
||||
try:
|
||||
# Split by "data: " to get individual SSE events
|
||||
events = re.split(b"data: ", chunk)
|
||||
for event_data in events:
|
||||
if (
|
||||
not event_data
|
||||
or event_data.strip() == b"[DONE]"
|
||||
or event_data.strip() == b""
|
||||
):
|
||||
continue
|
||||
|
||||
try:
|
||||
data = json.loads(event_data)
|
||||
if (
|
||||
"usage" in data
|
||||
and data["usage"] is not None
|
||||
and isinstance(data["usage"], dict)
|
||||
):
|
||||
# Found usage data, calculate cost
|
||||
# Create a new session for this operation
|
||||
async with create_session() as new_session:
|
||||
# Re-fetch the key in the new session
|
||||
fresh_key = await new_session.get(
|
||||
key.__class__, key.hashed_key
|
||||
)
|
||||
if fresh_key:
|
||||
cost_data = (
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, data, new_session
|
||||
)
|
||||
)
|
||||
# Format as SSE and yield
|
||||
cost_json = json.dumps(
|
||||
{"cost": cost_data}
|
||||
)
|
||||
yield f"data: {cost_json}\n\n".encode()
|
||||
break
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error processing streaming response for cost: {e}")
|
||||
|
||||
result = await handle_streaming_chat_completion(response, key, session)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
background_tasks.add_task(client.aclose)
|
||||
return StreamingResponse(
|
||||
stream_with_cost(),
|
||||
status_code=response.status_code,
|
||||
headers=dict(response.headers),
|
||||
background=background_tasks,
|
||||
)
|
||||
result.background = background_tasks
|
||||
return result
|
||||
|
||||
elif response.status_code == 200 and "application/json" in content_type:
|
||||
# Handle non-streaming response
|
||||
try:
|
||||
content = await response.aread()
|
||||
response_json = json.loads(content)
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
key, response_json, session
|
||||
return await handle_non_streaming_chat_completion(
|
||||
response, key, session
|
||||
)
|
||||
response_json["cost"] = cost_data
|
||||
|
||||
response_headers = dict(response.headers)
|
||||
|
||||
# Remove Transfer-Encoding header to avoid conflict with Content-Length header in common nginx setups
|
||||
if "transfer-encoding" in response_headers:
|
||||
del response_headers["transfer-encoding"]
|
||||
|
||||
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:
|
||||
print(f"Failed to parse JSON from upstream response: {e}")
|
||||
except Exception as e:
|
||||
print(f"Error adjusting payment for tokens: {e}")
|
||||
finally:
|
||||
await response.aclose()
|
||||
await client.aclose()
|
||||
@@ -250,7 +278,6 @@ async def proxy(
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
background_tasks.add_task(client.aclose)
|
||||
background_tasks.add_task(pay_out)
|
||||
|
||||
return StreamingResponse(
|
||||
response.aiter_bytes(),
|
||||
@@ -279,19 +306,8 @@ async def proxy(
|
||||
else:
|
||||
error_message = f"Error connecting to upstream service: {error_type}"
|
||||
|
||||
return Response(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": error_message,
|
||||
"type": "upstream_error",
|
||||
"code": 502,
|
||||
}
|
||||
}
|
||||
),
|
||||
status_code=502,
|
||||
media_type="application/json",
|
||||
)
|
||||
return create_error_response("upstream_error", error_message, 502)
|
||||
|
||||
except Exception as exc:
|
||||
await client.aclose()
|
||||
import traceback
|
||||
@@ -303,16 +319,78 @@ async def proxy(
|
||||
f"path={path}, query_params={dict(request.query_params)}\n"
|
||||
f"Traceback:\n{tb}"
|
||||
)
|
||||
return Response(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": "An unexpected server error occurred",
|
||||
"type": "internal_error",
|
||||
"code": 500,
|
||||
}
|
||||
}
|
||||
),
|
||||
status_code=500,
|
||||
media_type="application/json",
|
||||
return create_error_response(
|
||||
"internal_error", "An unexpected server error occurred", 500
|
||||
)
|
||||
|
||||
|
||||
@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:
|
||||
# Check for X-Cashu header first
|
||||
if x_cashu := request.headers.get("X-Cashu", None):
|
||||
key = await validate_bearer_key(
|
||||
sha256(x_cashu.encode()).hexdigest(), session, "X-CASHU"
|
||||
)
|
||||
|
||||
if auth := request.headers.get("Authorization", None):
|
||||
key = await get_bearer_token_key(request, path, session, auth)
|
||||
|
||||
else:
|
||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||
|
||||
request_body, error_response = await validate_and_read_json_body(request, path)
|
||||
if error_response:
|
||||
return error_response
|
||||
|
||||
await pay_for_request(key, session, request, request_body)
|
||||
|
||||
# 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
|
||||
)
|
||||
|
||||
if key.refund_address == "X-CASHU":
|
||||
refund_token = await x_cashu_refund(key, session)
|
||||
response.headers["X-Cashu-Refund"] = refund_token
|
||||
|
||||
return response
|
||||
|
||||
|
||||
async def get_bearer_token_key(
|
||||
request: Request, path: str, session: AsyncSession, auth: str
|
||||
) -> ApiKey:
|
||||
"""Handle bearer token authentication proxy requests."""
|
||||
# Handle regular bearer token authentication
|
||||
auth = request.headers.get("Authorization", "")
|
||||
bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else ""
|
||||
refund_address = request.headers.get("Refund-LNURL", None)
|
||||
key_expiry_time = request.headers.get("Key-Expiry-Time", None)
|
||||
|
||||
# Validate key_expiry_time header
|
||||
if key_expiry_time:
|
||||
try:
|
||||
key_expiry_time = int(key_expiry_time) # type: ignore
|
||||
except ValueError:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Invalid Key-Expiry-Time: must be a valid Unix timestamp",
|
||||
)
|
||||
if not refund_address:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Error: Refund-LNURL header required when using Key-Expiry-Time",
|
||||
)
|
||||
else:
|
||||
key_expiry_time = None
|
||||
|
||||
return await validate_bearer_key(
|
||||
bearer_key,
|
||||
session,
|
||||
refund_address,
|
||||
key_expiry_time, # type: ignore
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user