mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
x cashu
This commit is contained in:
+74
-37
@@ -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",
|
||||
}
|
||||
|
||||
+59
-68
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user