mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
currency unit passthrough
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user