x_cashu mvp + proxy refactor

This commit is contained in:
Shroominic
2025-06-28 14:17:09 -03:00
parent 1959e3a0fe
commit 61528bea96
2 changed files with 272 additions and 186 deletions
+8
View File
@@ -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
View File
@@ -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
)