diff --git a/router/account.py b/router/account.py index e87099d8..f9b4e6b9 100644 --- a/router/account.py +++ b/router/account.py @@ -76,6 +76,7 @@ async def refund_wallet_endpoint( status_code=400, detail="Balance too small to refund (less than 1 sat)" ) + # TODO: choose currency and mint based on what user has configured token = await wallet().send(remaining_balance_sats) result = {"msats": remaining_balance_msats, "recipient": None, "token": token} diff --git a/router/auth.py b/router/auth.py index bca663ec..1cb93dd0 100644 --- a/router/auth.py +++ b/router/auth.py @@ -6,10 +6,7 @@ from sqlmodel import col, update from .cashu import credit_balance from .db import ApiKey, AsyncSession -from .models import MODELS from .payment.cost_caculation import ( - COST_PER_REQUEST, - MODEL_BASED_PRICING, CostData, CostDataError, MaxCostData, @@ -101,15 +98,8 @@ async def validate_bearer_key( ) -async def pay_for_request( - key: ApiKey, - session: AsyncSession, - 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: - cost_per_request = get_max_cost_for_model(model=body["model"]) +async def pay_for_request(key: ApiKey, session: AsyncSession, body: dict) -> None: + cost_per_request = get_max_cost_for_model(model=body["model"]) if key.balance < cost_per_request: raise HTTPException( diff --git a/router/cashu.py b/router/cashu.py index 6cd43682..7bf8518a 100644 --- a/router/cashu.py +++ b/router/cashu.py @@ -4,7 +4,7 @@ import time from typing import cast from sixty_nuts import Wallet -from sixty_nuts.mint import CurrencyUnit +from sixty_nuts.types import CurrencyUnit from sqlmodel import col, func, select, update from .db import ApiKey, AsyncSession, get_session @@ -97,16 +97,19 @@ async def periodic_payout() -> None: async def credit_balance(cashu_token: str, key: ApiKey, session: AsyncSession) -> int: """Redeem a Cashu token and credit the amount to the API key balance.""" try: - amount_sats, _ = await wallet().redeem(cashu_token) + amount, unit = await wallet().redeem(cashu_token) except Exception as e: print(f"Error in credit_balance: {e}") # Ensure the balance cannot become negative if redeem fails return 0 - if amount_sats <= 0: + if amount <= 0: return 0 - amount_msats = amount_sats * 1000 + if unit == "msat": + amount_msats = amount + else: + amount_msats = amount * 1000 # Apply the balance change atomically to avoid race conditions when topping # up the same key concurrently. @@ -192,14 +195,15 @@ async def refund_balance(amount_msats: int, key: ApiKey, session: AsyncSession) return await wallet().send_to_lnurl(key.refund_address, amount=amount_sats) -async def x_cashu_refund(key: ApiKey, session: AsyncSession) -> str: - refund_token = await wallet().send(key.balance) +async def x_cashu_refund(key: ApiKey, session: AsyncSession, unit: CurrencyUnit) -> str: + refund_token = await wallet().send(key.balance, unit=unit) await session.delete(key) await session.commit() return refund_token async def redeem(cashu_token: str, lnurl: str) -> int: - amount_sats, _ = await wallet().redeem(cashu_token) - await wallet().send_to_lnurl(lnurl, amount=amount_sats) - return amount_sats + amount, unit = await wallet().redeem(cashu_token) + unit = cast(CurrencyUnit, unit) + await wallet().send_to_lnurl(lnurl, amount=amount, unit=unit) + return amount diff --git a/router/payment/helpers.py b/router/payment/helpers.py index efe4aa4e..b02da934 100644 --- a/router/payment/helpers.py +++ b/router/payment/helpers.py @@ -1,37 +1,39 @@ import base64 import json import os -from typing import Literal import cbor2 from fastapi import HTTPException, Response +from sixty_nuts.types import CurrencyUnit from router.models import MODELS -from router.payment.cost_caculation import COST_PER_REQUEST +from router.payment.cost_caculation import COST_PER_REQUEST, MODEL_BASED_PRICING UPSTREAM_BASE_URL = os.environ["UPSTREAM_BASE_URL"] UPSTREAM_API_KEY = os.environ.get("UPSTREAM_API_KEY", "") -def check_token_balance( - headers: dict, body: dict, unit: Literal["sat", "msat"] -) -> None: +def get_cost_per_request(model: str | None = None) -> int: + if MODEL_BASED_PRICING and MODELS and model: + return get_max_cost_for_model(model=model) + return COST_PER_REQUEST + + +def check_token_balance(headers: dict, body: dict) -> CurrencyUnit: 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"]) + cost = get_cost_per_request(model=body.get("model", None)) 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 == "msat": - pass - elif unit == "sat": + unit: CurrencyUnit = _token["unit"] + if unit == "sat": amount *= 1000 - if amount < COST_PER_REQUEST: + if amount < cost: raise HTTPException(status_code=413, detail="Insufficient balance") elif cashu_token.startswith("cashuB"): _token = base64_token_cbor(cashu_token) @@ -39,10 +41,11 @@ def check_token_balance( unit = _token["u"] if unit == "sat": amount *= 1000 - if amount < COST_PER_REQUEST: + if amount < cost: raise HTTPException(status_code=413, detail="Insufficient balance") else: raise HTTPException(status_code=401, detail="Unauthorized") + return unit def base64_token_json(cashu_token: str) -> dict: @@ -66,6 +69,8 @@ def base64_token_cbor(cashu_token: str) -> dict: def get_max_cost_for_model(model: str) -> int: + if not MODEL_BASED_PRICING or not MODELS: + return COST_PER_REQUEST if model not in [model.id for model in MODELS]: return COST_PER_REQUEST for m in MODELS: diff --git a/router/payment/x_cashu.py b/router/payment/x_cashu.py index f3cb5b47..3ef046fb 100644 --- a/router/payment/x_cashu.py +++ b/router/payment/x_cashu.py @@ -5,6 +5,7 @@ from typing import AsyncGenerator, Literal, cast import httpx from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse +from sixty_nuts.types import CurrencyUnit from router.cashu import wallet from router.payment.cost_caculation import ( @@ -25,13 +26,13 @@ async def x_cashu_handler( request: Request, x_cashu_token: str, path: str ) -> Response | StreamingResponse: headers = dict(request.headers) - amount, _ = await redeem_token(x_cashu_token) + amount, unit = await redeem_token(x_cashu_token) headers = prepare_upstream_headers(dict(request.headers)) - return await forward_to_upstream(request, path, headers, amount) + return await forward_to_upstream(request, path, headers, amount, unit) async def forward_to_upstream( - request: Request, path: str, headers: dict, amount: int + request: Request, path: str, headers: dict, amount: int, unit: CurrencyUnit ) -> Response | StreamingResponse: """Forward request to upstream and handle the response.""" if path.startswith("v1/"): @@ -55,7 +56,7 @@ async def forward_to_upstream( ) if path.endswith("chat/completions"): - result = await handle_x_cashu_chat_completion(response, amount) + result = await handle_x_cashu_chat_completion(response, amount, unit) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) result.background = background_tasks @@ -85,7 +86,7 @@ async def forward_to_upstream( async def handle_x_cashu_chat_completion( - response: httpx.Response, amount: int + response: httpx.Response, amount: int, unit: CurrencyUnit ) -> StreamingResponse | Response: """Handle both streaming and non-streaming chat completion responses with token-based pricing.""" try: @@ -95,10 +96,12 @@ async def handle_x_cashu_chat_completion( if is_streaming: print("Detected streaming response, processing SSE format") - return await handle_streaming_response(content_str, response, amount) + return await handle_streaming_response(content_str, response, amount, unit) else: print("Detected non-streaming response, processing as JSON") - return await handle_non_streaming_response(content_str, response, amount) + return await handle_non_streaming_response( + content_str, response, amount, unit + ) except Exception as e: print(f"Error processing chat completion response: {e}") @@ -111,7 +114,7 @@ async def handle_x_cashu_chat_completion( async def handle_streaming_response( - content_str: str, response: httpx.Response, amount: int + content_str: str, response: httpx.Response, amount: int, unit: CurrencyUnit ) -> StreamingResponse: """Handle Server-Sent Events (SSE) streaming response.""" # For streaming responses, we'll extract the final usage data @@ -144,7 +147,7 @@ async def handle_streaming_response( if cost_data: refund_amount = amount - cost_data.total_msats if refund_amount > 0: - refund_token = await send_refund(refund_amount) + refund_token = await send_refund(refund_amount, unit) response.headers["X-Cashu"] = refund_token print(f"Refunded {refund_amount} msats") except Exception as e: @@ -169,7 +172,7 @@ async def handle_streaming_response( async def handle_non_streaming_response( - content_str: str, response: httpx.Response, amount: int + content_str: str, response: httpx.Response, amount: int, unit: CurrencyUnit ) -> Response: """Handle regular JSON response.""" try: @@ -201,7 +204,7 @@ async def handle_non_streaming_response( refund_amount = amount - cost_data.total_msats print("refund: ", refund_amount) if refund_amount > 0: - refund_token = await send_refund(refund_amount) + refund_token = await send_refund(refund_amount, unit) response.headers["X-Cashu"] = refund_token print(f"Refunded {refund_amount} msats") @@ -266,9 +269,9 @@ async def redeem_token(x_cashu_token: str) -> tuple[int, Literal["sat", "msat"]] ) -async def send_refund(amount: int) -> str: +async def send_refund(amount: int, unit: CurrencyUnit, mint: str | None = None) -> str: try: - return await wallet().send(amount) + return await wallet().send(amount, unit=unit, mint_url=mint) except Exception as e: raise HTTPException( status_code=401, diff --git a/router/proxy.py b/router/proxy.py index 29239a69..e5087de6 100644 --- a/router/proxy.py +++ b/router/proxy.py @@ -263,12 +263,9 @@ async def proxy( status_code=400, media_type="application/json", ) - + unit = check_token_balance(headers, request_body_dict) # 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, "msat") return await x_cashu_handler(request, x_cashu, path) elif auth := headers.get("authorization", None): @@ -299,7 +296,7 @@ async def proxy( ) if response.status_code != 200 and key.refund_address == "X-CASHU": - refund_token = await x_cashu_refund(key, session) + refund_token = await x_cashu_refund(key, session, unit) response = Response( content=json.dumps( { @@ -318,7 +315,7 @@ async def proxy( return response if key.refund_address == "X-CASHU": - refund_token = await x_cashu_refund(key, session) + refund_token = await x_cashu_refund(key, session, unit) response.headers["X-Cashu"] = refund_token return response