currency unit passthrough

This commit is contained in:
Shroominic
2025-07-15 12:05:32 -03:00
parent 8c6ba30d26
commit 6bc1f2182f
6 changed files with 52 additions and 52 deletions
+1
View File
@@ -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}
+2 -12
View File
@@ -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(
+13 -9
View File
@@ -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
+17 -12
View File
@@ -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:
+16 -13
View File
@@ -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,
+3 -6
View File
@@ -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