From f5d5e0a3c9e23bbb37c01301754ede96c6fcd0c2 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 28 Jun 2025 09:37:15 -0300 Subject: [PATCH] x cashu --- router/auth.py | 111 +++++++++++++++++++++++++------------- router/proxy.py | 127 ++++++++++++++++++++------------------------ tests/test_proxy.py | 11 ++++ 3 files changed, 144 insertions(+), 105 deletions(-) diff --git a/router/auth.py b/router/auth.py index 1509ca10..4bdceff1 100644 --- a/router/auth.py +++ b/router/auth.py @@ -3,7 +3,7 @@ import json import os from typing import Optional -from fastapi import HTTPException, Request +from fastapi import HTTPException from sqlmodel import col, update from .cashu import credit_balance @@ -113,48 +113,85 @@ async def validate_bearer_key( ) +def base64_token_json(cashu_token: str) -> dict: + import base64 + + # Version 3 - JSON format + encoded = cashu_token[6:] # Remove "cashuA" + # Add correct padding – (-len) % 4 equals 0,1,2,3 + encoded += "=" * ((-len(encoded)) % 4) + + decoded = base64.urlsafe_b64decode(encoded).decode() + token_data = json.loads(decoded) + + return token_data + + +def base64_token_cbor(cashu_token: str) -> dict: + import base64 + + import cbor2 + + encoded = cashu_token[6:] # Remove "cashuB" + encoded += "=" * ((-len(encoded)) % 4) + decoded_bytes = base64.urlsafe_b64decode(encoded) + token_data = cbor2.loads(decoded_bytes) + return token_data + + +def check_token_balance(headers: dict, body: dict) -> None: + if x_cashu := headers.get("x-cashu", None): + cashu_token = x_cashu + elif auth := headers.get("authorization", None): + cashu_token = auth.split(" ")[1] + else: + raise HTTPException(status_code=401, detail="Unauthorized") + COST_PER_REQUEST = get_max_cost_for_model(model=body["model"]) + if cashu_token.startswith("cashuA"): + _token = base64_token_json(cashu_token) + amount = sum(p["amount"] for t in _token["token"] for p in t["proofs"]) + unit = _token["unit"] + if unit == "sat": + amount *= 1000 + if amount < COST_PER_REQUEST: + raise HTTPException(status_code=413, detail="Insufficient balance") + elif cashu_token.startswith("cashuB"): + _token = base64_token_cbor(cashu_token) + amount = sum(p["a"] for t in _token["t"] for p in t["p"]) + unit = _token["u"] + if unit == "sat": + amount *= 1000 + if amount < COST_PER_REQUEST: + raise HTTPException(status_code=413, detail="Insufficient balance") + else: + raise HTTPException(status_code=401, detail="Unauthorized") + + +def get_max_cost_for_model(model: str) -> int: + if model not in [model.id for model in MODELS]: + return COST_PER_REQUEST + for m in MODELS: + if m.id == model: + return m.sats_pricing.max_cost * 1000 # type: ignore + return COST_PER_REQUEST + + async def pay_for_request( key: ApiKey, session: AsyncSession, - request: Request | None, - request_body: bytes | None = None, + body: dict, ) -> None: + # Use global COST_PER_REQUEST as default, override if model-based pricing is enabled + cost_per_request = COST_PER_REQUEST if MODEL_BASED_PRICING and MODELS: - if request_body: - body = json.loads(request_body) - else: - body = await request.json() # type: ignore - if request_model := body.get("model"): - if request_model not in [model.id for model in MODELS]: - raise HTTPException( - status_code=400, - detail={ - "error": { - "message": f"Invalid model: {request_model}", - "type": "invalid_request_error", - "code": "model_not_found", - } - }, - ) - model = next(model for model in MODELS if model.id == request_model) - if key.balance < model.sats_pricing.max_cost * 1000: # type: ignore - raise HTTPException( - status_code=413, - detail={ - "error": { - "message": f"This model requires a minimum balance of {model.sats_pricing.max_cost} sats", # type: ignore - "type": "insufficient_quota", - "code": "insufficient_balance", - } - }, - ) + cost_per_request = get_max_cost_for_model(model=body["model"]) - if key.balance < COST_PER_REQUEST: + if key.balance < cost_per_request: raise HTTPException( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {COST_PER_REQUEST} mSats required. {key.balance} available.", + "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.", "type": "insufficient_quota", "code": "insufficient_balance", } @@ -165,10 +202,10 @@ async def pay_for_request( stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.balance) >= COST_PER_REQUEST) + .where(col(ApiKey.balance) >= cost_per_request) .values( - balance=col(ApiKey.balance) - COST_PER_REQUEST, - total_spent=col(ApiKey.total_spent) + COST_PER_REQUEST, + balance=col(ApiKey.balance) - cost_per_request, + total_spent=col(ApiKey.total_spent) + cost_per_request, total_requests=col(ApiKey.total_requests) + 1, ) ) @@ -180,7 +217,7 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {COST_PER_REQUEST} mSats required. {key.balance} available.", + "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.", "type": "insufficient_quota", "code": "insufficient_balance", } diff --git a/router/proxy.py b/router/proxy.py index f899ea06..bf1644ed 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -7,7 +7,12 @@ import httpx 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 .auth import ( + adjust_payment_for_tokens, + check_token_balance, + pay_for_request, + validate_bearer_key, +) from .cashu import x_cashu_refund from .db import ApiKey, AsyncSession, create_session, get_session @@ -17,57 +22,6 @@ UPSTREAM_API_KEY = os.environ.get("UPSTREAM_API_KEY", "") proxy_router = APIRouter() -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. - - 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 - - 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 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", - ) - - def prepare_upstream_headers(request_headers: dict) -> dict: """Prepare headers for upstream request, removing sensitive/problematic ones.""" headers = dict(request_headers) @@ -331,23 +285,43 @@ async def forward_to_upstream( async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) ) -> Response | StreamingResponse: - # todo check token balance without claiming and raise 402 or 413 if not enough - # check_token_balance(request, session) + request_body = await request.body() + headers = dict(request.headers) - if x_cashu := request.headers.get("X-Cashu", None): + # Parse JSON body if present, handle empty/invalid JSON + request_body_dict = {} + if request_body: + try: + request_body_dict = json.loads(request_body) + except json.JSONDecodeError: + return Response( + content=json.dumps( + {"error": {"type": "invalid_request_error", "code": "invalid_json"}} + ), + status_code=400, + media_type="application/json", + ) + + # Handle authentication + if x_cashu := headers.get("x-cashu", None): + # Check token balance before authentication for cashu tokens + if request_body_dict: + check_token_balance(headers, request_body_dict) key = await validate_bearer_key(x_cashu, session, "X-CASHU") - elif auth := request.headers.get("Authorization", None): - key = await get_bearer_token_key(request, path, session, auth) + elif auth := headers.get("authorization", None): + key = await get_bearer_token_key(headers, path, session, auth) else: - raise HTTPException(status_code=401, detail="Unauthorized") + return Response( + content=json.dumps({"detail": "Unauthorized"}), + status_code=401, + media_type="application/json", + ) - 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) + # Only pay for request if we have request body data (for completions endpoints) + if request_body_dict: + await pay_for_request(key, session, request_body_dict) # Prepare headers for upstream headers = prepare_upstream_headers(dict(request.headers)) @@ -357,6 +331,25 @@ async def proxy( request, path, headers, request_body, key, session ) + if response.status_code != 200 and key.refund_address == "X-CASHU": + refund_token = await x_cashu_refund(key, session) + response = Response( + content=json.dumps( + { + "error": { + "message": "Error forwarding request to upstream", + "type": "upstream_error", + "code": response.status_code, + "refund_token": refund_token, + } + } + ), + status_code=response.status_code, + media_type="application/json", + ) + response.headers["X-Cashu"] = refund_token + return response + if key.refund_address == "X-CASHU": refund_token = await x_cashu_refund(key, session) response.headers["X-Cashu"] = refund_token @@ -365,14 +358,12 @@ async def proxy( async def get_bearer_token_key( - request: Request, path: str, session: AsyncSession, auth: str + headers: dict, 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) + refund_address = headers.get("Refund-LNURL", None) + key_expiry_time = headers.get("Key-Expiry-Time", None) # Validate key_expiry_time header if key_expiry_time: diff --git a/tests/test_proxy.py b/tests/test_proxy.py index a07b18e1..af2df006 100644 --- a/tests/test_proxy.py +++ b/tests/test_proxy.py @@ -33,6 +33,17 @@ async def test_proxy_requires_authentication(async_client: AsyncClient) -> None: """Test that proxy endpoints require authentication.""" response = await async_client.post("/v1/chat/completions") + assert response.status_code == 401 + assert response.json()["detail"] == "Unauthorized" + + +@pytest.mark.asyncio +async def test_proxy_empty_bearer_token(async_client: AsyncClient) -> None: + """Test that proxy endpoints return structured error for empty bearer token.""" + response = await async_client.post( + "/v1/chat/completions", headers={"Authorization": "Bearer "} + ) + assert response.status_code == 401 assert ( "API key or Cashu token required"