Merge branch 'main' into kwsantiago/62-comprehensive-tests

This commit is contained in:
Shroominic
2025-08-03 15:55:43 -03:00
33 changed files with 3053 additions and 2892 deletions
+17 -25
View File
@@ -1,48 +1,40 @@
NAME = "Your Routstr Proxy Name"
DESCRIPTION = "A short Description"
# Any openai-compatible api endpoint
UPSTREAM_BASE_URL="https://api.openai.com/v1"
UPSTREAM_API_KEY="sk-21212121212121212121212121212121"
# UPSTREAM_PROVIDER_FEE=1 # 1 = no fees, 1.05 = 5% fees
# Lightning address used to receive funds
RECEIVE_LN_ADDRESS="shroominic@walletofsatoshi.com"
# A prepaid key you can use when you are using it yourself or while testing. Api-key doesn't need to be in any specific format but it should start with "sk-""
PREPAID_API_KEY="sk-" # Add any string/hash here. Replace with a new string for a new API
PREPAID_BALANCE="10000" # 10k sats
# When your cashu balance reaches this number of sats, send the funds to RECEIVE_LN_ADDRESS.
MINIMUM_PAYOUT = "100"
# Costs in Sats, if MODEL_BASED_PRICING is set to false
COST_PER_REQUEST="10"
COST_PER_1K_INPUT_TOKENS = "0"
COST_PER_1K_OUTPUT_TOKENS = "0"
# If set to true, pricing is loaded from the file specified by MODELS_PATH
# Defaults to "models.json" and falls back to "models.example.json" if missing
MODEL_BASED_PRICING = "false"
MODEL_BASED_PRICING = "true"
# MODELS_PATH="models.json"
# Time in seconds between each automatically refunding funds to users whose API keys have expired
# Setting this to "0" disables automatic refunds
REFUND_PROCESSING_INTERVAL = "3600"
# Costs in Sats, if MODEL_BASED_PRICING is set to false
# COST_PER_REQUEST="10"
# COST_PER_1K_INPUT_TOKENS = "0"
# COST_PER_1K_OUTPUT_TOKENS = "0"
# EXCHANGE_FEE = "1.005" # 0.5 % currency exchange fee
# password used to log into admin interface
ADMIN_PASSWORD="XXX"
ADMIN_PASSWORD="CHANGE-THIS"
# NPUB of Nostr account
NSEC=""
NPUB="npub..."
# Public Endpoint
HTTP_URL="https://your.domain.com"
# Not used currently
HTTP_URL=""
# Tor Endpoint (copy from docker logs)
ONION_URL=".onion"
# Not used currently
ONION_URL="XXX.onion"
RELAYS="wss://relay.damus.io,wss://relay.nostr.band"
RELAYS="wss://relay.routstr.com,wss://relay.nostr.band"
CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org"
# Development
# DEBUG=TRUE
# LOG_LEVEL=TRACE
-5
View File
@@ -1,5 +0,0 @@
- test if currency and payment amount is correct
- test payout
- make tor work
-
+2 -1
View File
@@ -13,7 +13,8 @@ RUN apk add git
COPY uv.lock pyproject.toml ./
RUN uv sync
RUN uv add git+https://github.com/saschanaz/secp256k1-py.git#branch=upgrade060
# RUN uv sync
WORKDIR /app
+4
View File
@@ -84,6 +84,10 @@ The most common settings are shown below. See `.env.example` for the full list.
- `REFUND_PROCESSING_INTERVAL` – Seconds between automatic refunds
- `ADMIN_PASSWORD` – Password for the `/admin` dashboard
## Withdrawing Balance
Go to `https://<your.routstr.proxy>/admin/` (NOTE: be sure to add the '/' at the end), enter the `ADMIN_PASSWORD` you set above and withdraw your balance as a Cashu token.
## Example Client
`example.py` shows how to use the proxy with the official OpenAI client:
+104 -14
View File
@@ -1,11 +1,12 @@
{
"models": [
{
"id": "google/gemini-2.5-pro-preview",
"id": "google/gemini-2.5-flash",
"canonical_slug": "google/gemini-2.5-flash",
"hugging_face_id": "",
"name": "Google: Gemini 2.5 Pro Preview 06-05",
"created": 1749137257,
"description": "Gemini 2.5 Pro is Google\u2019s state-of-the-art AI model designed for advanced reasoning, coding, mathematics, and scientific tasks. It employs \u201cthinking\u201d capabilities, enabling it to reason through responses with enhanced accuracy and nuanced context handling. Gemini 2.5 Pro achieves top-tier performance on multiple benchmarks, including first-place positioning on the LMArena leaderboard, reflecting superior human-preference alignment and complex problem-solving abilities.\n",
"name": "Google: Gemini 2.5 Flash",
"created": 1750172488,
"description": "Gemini 2.5 Flash is Google's state-of-the-art workhorse model, specifically designed for advanced reasoning, coding, mathematics, and scientific tasks. It includes built-in \"thinking\" capabilities, enabling it to provide responses with greater accuracy and nuanced context handling. \n\nAdditionally, Gemini 2.5 Flash is configurable through the \"max tokens for reasoning\" parameter, as described in the documentation (https://openrouter.ai/docs/use-cases/reasoning-tokens#max-tokens-for-reasoning).",
"context_length": 1048576,
"architecture": {
"modality": "text+image->text",
@@ -21,18 +22,107 @@
"instruct_type": null
},
"pricing": {
"prompt": "0.00000125",
"completion": "0.00001",
"prompt": "0.0000003",
"completion": "0.0000025",
"request": "0",
"image": "0.00516",
"image": "0.001238",
"web_search": "0",
"internal_reasoning": "0",
"input_cache_read": "0.00000031",
"input_cache_write": "0.000001625"
"input_cache_read": "0.000000075",
"input_cache_write": "0.0000003833"
},
"top_provider": {
"context_length": 1048576,
"max_completion_tokens": 65536,
"max_completion_tokens": 65535,
"is_moderated": false
},
"per_request_limits": null,
"supported_parameters": [
"max_tokens",
"temperature",
"top_p",
"tools",
"tool_choice",
"stop",
"response_format",
"structured_outputs"
]
},
{
"id": "openai/o3-pro",
"canonical_slug": "openai/o3-pro-2025-06-10",
"hugging_face_id": "",
"name": "OpenAI: o3 Pro",
"created": 1749598352,
"description": "The o-series of models are trained with reinforcement learning to think before they answer and perform complex reasoning. The o3-pro model uses more compute to think harder and provide consistently better answers.\n\nNote that BYOK is required for this model. Set up here: https://openrouter.ai/settings/integrations",
"context_length": 200000,
"architecture": {
"modality": "text+image->text",
"input_modalities": [
"text",
"file",
"image"
],
"output_modalities": [
"text"
],
"tokenizer": "Other",
"instruct_type": null
},
"pricing": {
"prompt": "0.00002",
"completion": "0.00008",
"request": "0",
"image": "0.0153",
"web_search": "0",
"internal_reasoning": "0"
},
"top_provider": {
"context_length": 200000,
"max_completion_tokens": 100000,
"is_moderated": true
},
"per_request_limits": null,
"supported_parameters": [
"tools",
"tool_choice",
"seed",
"max_tokens",
"response_format",
"structured_outputs"
]
},
{
"id": "x-ai/grok-3-mini",
"canonical_slug": "x-ai/grok-3-mini",
"hugging_face_id": "",
"name": "xAI: Grok 3 Mini",
"created": 1749583245,
"description": "A lightweight model that thinks before responding. Fast, smart, and great for logic-based tasks that do not require deep domain knowledge. The raw thinking traces are accessible.",
"context_length": 131072,
"architecture": {
"modality": "text->text",
"input_modalities": [
"text"
],
"output_modalities": [
"text"
],
"tokenizer": "Grok",
"instruct_type": null
},
"pricing": {
"prompt": "0.0000003",
"completion": "0.0000005",
"request": "0",
"image": "0",
"web_search": "0",
"internal_reasoning": "0",
"input_cache_read": "0.000000075"
},
"top_provider": {
"context_length": 131072,
"max_completion_tokens": null,
"is_moderated": false
},
"per_request_limits": null,
@@ -45,11 +135,11 @@
"reasoning",
"include_reasoning",
"structured_outputs",
"response_format",
"stop",
"frequency_penalty",
"presence_penalty",
"seed"
"seed",
"logprobs",
"top_logprobs",
"response_format"
]
}
]
+7 -2
View File
@@ -1,17 +1,19 @@
[project]
name = "routstr"
version = "0.0.1"
version = "0.1.0"
description = "Payment proxy for your LLM endpoint using cashu and nostr."
readme = "README.md"
requires-python = ">=3.11"
dependencies = [
"fastapi[standard]>=0.115",
"aiosqlite>=0.20",
"sixty-nuts>=0.1.4",
"sqlmodel>=0.0.24",
"httpx[socks]>=0.25.2",
"greenlet>=3.2.1",
"python-json-logger>=2.0.0",
"cashu",
"secp256k1",
"marshmallow>=3.13,<4.0",
]
[dependency-groups]
@@ -63,3 +65,6 @@ check_untyped_defs = true
disallow_untyped_calls = true
disallow_incomplete_defs = true
disallow_untyped_decorators = true
[tool.uv.sources]
secp256k1 = { git = "https://github.com/saschanaz/secp256k1-py", branch = "upgrade060" }
+1 -1
View File
@@ -2,6 +2,6 @@ import dotenv
dotenv.load_dotenv()
from .main import app as fastapi_app # noqa
from .core.main import app as fastapi_app # noqa
__all__ = ["fastapi_app"]
-162
View File
@@ -1,162 +0,0 @@
import os
from datetime import datetime, timezone
from fastapi import APIRouter, Request
from fastapi.responses import HTMLResponse
from sqlmodel import select
from .cashu import wallet
from .db import ApiKey, create_session
admin_router = APIRouter(prefix="/admin")
def login_form() -> str:
return """<!DOCTYPE html>
<html>
<head>
<style>
body {
font-family: Arial, sans-serif;
display: flex;
justify-content: center;
align-items: center;
height: 100vh;
margin: 0;
}
form {
display: flex;
flex-direction: column;
gap: 10px;
}
input[type="password"] {
padding: 8px;
}
button {
padding: 8px;
cursor: pointer;
}
</style>
<script>
function handleSubmit(e) {
e.preventDefault();
const password = document.getElementById('password').value;
document.cookie = `admin_password=${password}; path=/; max-age=86400`;
window.location.reload();
}
</script>
</head>
<body>
<form onsubmit="handleSubmit(event)">
<input type="password" id="password" placeholder="Admin Password" required>
<button type="submit">Login</button>
</form>
</body>
</html>
"""
def info(content: str) -> str:
return f"""<!DOCTYPE html>
<html>
<head>
<style>
body {{
font-family: Arial, sans-serif;
display: flex;
justify-content: center;
align-items: center;
height: 100vh;
margin: 0;
}}
</style>
</head>
<body>
<div style="text-align: center;">
{content}
</div>
</body>
</html>
"""
def admin_auth() -> str:
if os.getenv("ADMIN_PASSWORD", "") == "":
return info("Please set a secure ADMIN_PASSWORD= in your ENV variables.")
else:
return login_form()
async def dashboard(request: Request) -> str:
# fetch cashu / api-key data from database
async with create_session() as session:
result = await session.exec(select(ApiKey))
api_keys = result.all()
api_keys_table_rows = []
for key in api_keys:
expiry_time_utc = (
datetime.fromtimestamp(key.key_expiry_time, tz=timezone.utc)
if key.key_expiry_time is not None
else None
)
expiry_time_human_readable = (
expiry_time_utc.strftime("%Y-%m-%d %H:%M:%S") if expiry_time_utc else ""
)
api_keys_table_rows.append(
f"<tr><td>{key.hashed_key}</td><td>{key.balance}</td><td>{key.total_spent}</td><td>{key.total_requests}</td><td>{key.refund_address}</td><td>{'{} ({} UTC)'.format(key.key_expiry_time, expiry_time_human_readable) if key.key_expiry_time else key.key_expiry_time}</td></tr>"
)
# Calculate the total balance of all API keys using integer arithmetic to
# avoid rounding issues.
total_user_balance = sum(key.balance for key in api_keys) // 1000
# Fetch balance from cashu
current_balance = await wallet().get_balance()
owner_balance = current_balance - total_user_balance
return f"""<!DOCTYPE html>
<html>
<head>
<style>
table {{
width: 100%;
border-collapse: collapse;
}}
th, td {{
border: 1px solid black;
padding: 8px;
text-align: left;
}}
</style>
</head>
<body>
<h1>Admin Dashboard</h1>
<h2>Current Cashu Balance</h2>
<p>Your Balance: {owner_balance} sats</p>
<p>The balance is calculated by subtracting the combined user balance from the total Cashu wallet balance.</p>
<p>Total Cashu Balance: {current_balance} sats</p>
<p>User Balance: {total_user_balance} sats</p>
<h2>User's API Keys</h2>
<table>
<tr>
<th>Hashed Key</th>
<th>Balance (mSats)</th>
<th>Total Spent (mSats)</th>
<th>Total Requests</th>
<th>Refund Address</th>
<th>Refund Time</th>
</tr>
{"".join(api_keys_table_rows)}
</table>
</body>
</html>
"""
@admin_router.get("/", response_class=HTMLResponse)
async def admin(request: Request) -> str:
admin_cookie = request.cookies.get("admin_password")
if admin_cookie and admin_cookie == os.getenv("ADMIN_PASSWORD"):
return await dashboard(request)
return admin_auth()
+3 -3
View File
@@ -4,9 +4,8 @@ from typing import Optional
from fastapi import HTTPException
from sqlmodel import col, update
from .cashu import credit_balance
from .db import ApiKey, AsyncSession
from .logging.logging_config import get_logger
from .core import get_logger
from .core.db import ApiKey, AsyncSession
from .payment.cost_caculation import (
CostData,
CostDataError,
@@ -14,6 +13,7 @@ from .payment.cost_caculation import (
calculate_cost,
)
from .payment.helpers import get_max_cost_for_model
from .wallet import credit_balance
logger = get_logger(__name__)
+22 -29
View File
@@ -3,15 +3,11 @@ from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException
from .auth import validate_bearer_key
from .cashu import (
credit_balance,
delete_key_if_zero_balance,
refund_balance,
wallet,
)
from .db import ApiKey, AsyncSession, get_session
from .core.db import ApiKey, AsyncSession, get_session
from .wallet import credit_balance, send_to_lnurl, send_token
wallet_router = APIRouter(prefix="/v1/wallet")
router = APIRouter()
balance_router = APIRouter(prefix="/v1/balance")
async def get_key_from_header(
@@ -28,7 +24,7 @@ async def get_key_from_header(
# TODO: remove this endpoint when frontend is updated
@wallet_router.get("/")
@router.get("/", include_in_schema=False)
async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
return {
"api_key": "sk-" + key.hashed_key,
@@ -36,7 +32,7 @@ async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
}
@wallet_router.get("/info")
@router.get("/info")
async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
return {
"api_key": "sk-" + key.hashed_key,
@@ -44,7 +40,7 @@ async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
}
@wallet_router.post("/topup")
@router.post("/topup")
async def topup_wallet_endpoint(
cashu_token: str,
key: ApiKey = Depends(get_key_from_header),
@@ -95,7 +91,7 @@ async def topup_wallet_endpoint(
return {"msats": amount_msats}
@wallet_router.post("/refund")
@router.post("/refund")
async def refund_wallet_endpoint(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
@@ -107,9 +103,8 @@ async def refund_wallet_endpoint(
# Perform refund operation first, before modifying balance
if key.refund_address:
# refund_balance handles balance update and key deletion
await refund_balance(remaining_balance_msats, key, session)
return {"recipient": key.refund_address, "msats": remaining_balance_msats}
await send_to_lnurl(remaining_balance_msats, "msat", key.refund_address)
result = {"recipient": key.refund_address, "msat": remaining_balance_msats}
else:
# Convert msats to sats for cashu wallet
remaining_balance_sats = remaining_balance_msats // 1000
@@ -119,24 +114,17 @@ async def refund_wallet_endpoint(
)
# TODO: choose currency and mint based on what user has configured
try:
token = await wallet().send(remaining_balance_sats)
except Exception as e:
# Handle mint service errors
raise HTTPException(
status_code=503, detail=f"Mint service unavailable: {str(e)}"
)
token = await send_token(remaining_balance_sats, "sat")
# Only for token refunds, we need to manually update balance and delete key
key.balance = 0
session.add(key)
await session.commit()
await delete_key_if_zero_balance(key, session)
result = {"recipient": None, "msat": remaining_balance_msats, "token": token}
return {"msats": remaining_balance_msats, "recipient": None, "token": token}
await session.delete(key)
await session.commit()
return result
@wallet_router.api_route(
@router.api_route(
"/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"],
include_in_schema=False,
@@ -146,3 +134,8 @@ async def wallet_catch_all(path: str) -> NoReturn:
raise HTTPException(
status_code=404, detail="Not found check /docs for available endpoints"
)
balance_router.include_router(router)
deprecated_wallet_router = APIRouter(prefix="/v1/wallet", include_in_schema=False)
deprecated_wallet_router.include_router(router)
-537
View File
@@ -1,537 +0,0 @@
import asyncio
import os
import time
from typing import cast
from sixty_nuts import Wallet
from sixty_nuts.types import CurrencyUnit
from sqlmodel import col, func, select, update
from .db import ApiKey, AsyncSession, get_session
from .logging.logging_config import get_logger
logger = get_logger(__name__)
RECEIVE_LN_ADDRESS = os.environ["RECEIVE_LN_ADDRESS"]
MINT = os.environ.get("MINT", "https://mint.minibits.cash/Bitcoin")
MINIMUM_PAYOUT = int(os.environ.get("MINIMUM_PAYOUT", 100))
REFUND_PROCESSING_INTERVAL = int(os.environ.get("REFUND_PROCESSING_INTERVAL", 3600))
PAYOUT_INTERVAL = int(os.environ.get("PAYOUT_INTERVAL", 300)) # Default 5 minutes
DEV_LN_ADDRESS = "routstr@minibits.cash"
DEVS_DONATION_RATE = float(os.environ.get("DEVS_DONATION_RATE", 0.021)) # 2.1%
NSEC = os.environ["NSEC"] # Nostr private key for the wallet
logger.info(
"Cashu module initialized",
extra={
"mint": MINT,
"minimum_payout": MINIMUM_PAYOUT,
"refund_processing_interval": REFUND_PROCESSING_INTERVAL,
"payout_interval": PAYOUT_INTERVAL,
"devs_donation_rate": DEVS_DONATION_RATE,
},
)
wallet_instance: Wallet | None = None
async def init_wallet() -> None:
"""Initialize the Cashu wallet."""
global wallet_instance
try:
logger.info("Initializing Cashu wallet", extra={"mint": MINT})
wallet_instance = await Wallet.create(nsec=NSEC)
logger.info("Cashu wallet initialized successfully")
except Exception as e:
logger.error(
"Failed to initialize Cashu wallet",
extra={"error": str(e), "error_type": type(e).__name__, "mint": MINT},
)
raise
def wallet() -> Wallet:
"""Get the wallet instance."""
global wallet_instance
if wallet_instance is None:
logger.error("Wallet not initialized - call init_wallet() first")
raise ValueError("Wallet not initialized")
return wallet_instance
async def delete_key_if_zero_balance(key: ApiKey, session: AsyncSession) -> None:
"""Delete the given API key if its balance is zero."""
if key.balance == 0:
logger.info(
"Deleting API key with zero balance",
extra={"key_hash": key.hashed_key[:8] + "...", "balance": key.balance},
)
await session.delete(key)
await session.commit()
async def pay_out() -> None:
"""
Calculates the pay-out amount based on the spent balance, profit, and donation rate.
"""
try:
logger.debug("Starting payout process")
from .db import create_session
async with create_session() as session:
result = await session.exec(
select(func.sum(col(ApiKey.balance))).where(ApiKey.balance > 0)
)
balance = result.one_or_none()
if not balance:
logger.debug("No balance to pay out")
return
user_balance_sats = balance // 1000
wallet_balance_sats = await wallet().get_balance()
logger.debug(
"Payout calculation",
extra={
"user_balance_sats": user_balance_sats,
"wallet_balance_sats": wallet_balance_sats,
},
)
# Handle edge cases more gracefully
if wallet_balance_sats < user_balance_sats:
logger.warning(
"Insufficient wallet balance for payout",
extra={
"wallet_balance_sats": wallet_balance_sats,
"user_balance_sats": user_balance_sats,
"shortfall_sats": user_balance_sats - wallet_balance_sats,
},
)
return
if (revenue := wallet_balance_sats - user_balance_sats) <= MINIMUM_PAYOUT:
logger.debug(
"Revenue below minimum payout threshold",
extra={"revenue_sats": revenue, "minimum_payout": MINIMUM_PAYOUT},
)
return
devs_donation = int(revenue * DEVS_DONATION_RATE)
owners_draw = revenue - devs_donation
logger.info(
"Processing payout",
extra={
"revenue_sats": revenue,
"devs_donation_sats": devs_donation,
"owners_draw_sats": owners_draw,
"donation_rate": DEVS_DONATION_RATE,
},
)
# Send payouts
try:
await wallet().send_to_lnurl(RECEIVE_LN_ADDRESS, owners_draw)
logger.info(
"Owner payout sent successfully",
extra={
"amount_sats": owners_draw,
"address": RECEIVE_LN_ADDRESS[:10] + "...",
},
)
await wallet().send_to_lnurl(DEV_LN_ADDRESS, devs_donation)
logger.info(
"Developer donation sent successfully",
extra={"amount_sats": devs_donation, "address": DEV_LN_ADDRESS},
)
except Exception as payout_error:
logger.error(
"Failed to send payouts",
extra={
"error": str(payout_error),
"error_type": type(payout_error).__name__,
"owners_draw_sats": owners_draw,
"devs_donation_sats": devs_donation,
},
)
raise
except Exception as e:
logger.error(
"Error in payout process",
extra={"error": str(e), "error_type": type(e).__name__},
)
# Periodic payout task
async def periodic_payout() -> None:
"""Periodically process payouts."""
logger.info("Starting periodic payout task", extra={"interval_seconds": 300})
while True:
try:
await asyncio.sleep(300) # Run every 5 minutes
await pay_out()
except asyncio.CancelledError:
logger.info("Periodic payout task cancelled")
break
except Exception as e:
logger.error(
"Error in periodic payout",
extra={"error": str(e), "error_type": type(e).__name__},
)
# Continue running even if payout fails
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."""
logger.debug(
"Starting token redemption",
extra={
"key_hash": key.hashed_key[:8] + "...",
"token_preview": cashu_token[:20] + "..."
if len(cashu_token) > 20
else cashu_token,
},
)
try:
amount, unit = await wallet().redeem(cashu_token)
logger.info(
"Token redeemed successfully",
extra={
"amount": amount,
"unit": unit,
"key_hash": key.hashed_key[:8] + "...",
},
)
except Exception as e:
logger.error(
"Token redemption failed",
extra={
"error": str(e),
"error_type": type(e).__name__,
"key_hash": key.hashed_key[:8] + "...",
"token_preview": cashu_token[:20] + "..."
if len(cashu_token) > 20
else cashu_token,
},
)
# Re-raise the exception so the caller can handle invalid tokens properly
raise ValueError(f"Token redemption failed: {str(e)}")
if amount <= 0:
logger.warning(
"Zero or negative amount redeemed",
extra={
"amount": amount,
"unit": unit,
"key_hash": key.hashed_key[:8] + "...",
},
)
return 0
if unit == "msat":
amount_msats = amount
else:
amount_msats = amount * 1000
logger.debug(
"Crediting balance",
extra={
"amount_msats": amount_msats,
"original_amount": amount,
"unit": unit,
"key_hash": key.hashed_key[:8] + "...",
},
)
# Apply the balance change atomically to avoid race conditions when topping
# up the same key concurrently.
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(balance=col(ApiKey.balance) + amount_msats)
)
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
await session.refresh(key)
logger.info(
"Balance credited successfully",
extra={
"credited_msats": amount_msats,
"new_balance_msats": key.balance,
"key_hash": key.hashed_key[:8] + "...",
},
)
return amount_msats
async def check_for_refunds() -> None:
"""
Periodically checks for API keys that are eligible for refunds and processes them.
Raises:
Exception: If an error occurs during the refund check process.
"""
# Setting REFUND_PROCESSING_INTERVAL to 0 disables it
if REFUND_PROCESSING_INTERVAL == 0:
logger.info("Automatic refund processing is disabled")
return
logger.info(
"Starting refund monitoring task",
extra={"interval_seconds": REFUND_PROCESSING_INTERVAL},
)
while True:
try:
logger.debug("Checking for expired keys requiring refunds")
async for session in get_session():
result = await session.exec(select(ApiKey))
keys = result.all()
current_time = int(time.time())
expired_keys = []
for key in keys:
if (
key.balance > 0
and key.refund_address
and key.key_expiry_time
and key.key_expiry_time < current_time
):
expired_keys.append(key)
if expired_keys:
logger.info(
"Found expired keys for refund",
extra={
"expired_count": len(expired_keys),
"current_time": current_time,
},
)
for key in expired_keys:
logger.info(
"Processing refund for expired key",
extra={
"key_hash": key.hashed_key[:8] + "...",
"balance_msats": key.balance,
"expiry_time": key.key_expiry_time,
"current_time": current_time,
"expired_seconds": current_time
- (key.key_expiry_time or 0),
},
)
try:
await refund_balance(key.balance, key, session)
await delete_key_if_zero_balance(key, session)
logger.info(
"Refund processed successfully",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
except Exception as refund_error:
logger.error(
"Failed to process refund",
extra={
"error": str(refund_error),
"error_type": type(refund_error).__name__,
"key_hash": key.hashed_key[:8] + "...",
"balance_msats": key.balance,
},
)
# Sleep for the specified interval before checking again
await asyncio.sleep(REFUND_PROCESSING_INTERVAL)
except asyncio.CancelledError:
logger.info("Refund monitoring task cancelled")
break
except Exception as e:
logger.error(
"Error during refund check",
extra={"error": str(e), "error_type": type(e).__name__},
)
async def refund_balance(amount_msats: int, key: ApiKey, session: AsyncSession) -> int:
"""Process a refund for an API key."""
if amount_msats <= 0:
amount_msats = key.balance
logger.info(
"Processing balance refund",
extra={
"amount_msats": amount_msats,
"key_hash": key.hashed_key[:8] + "...",
"refund_address": key.refund_address[:20] + "..."
if key.refund_address and len(key.refund_address) > 20
else key.refund_address,
},
)
# Convert msats to sats for cashu wallet
amount_sats = amount_msats // 1000
if amount_sats == 0:
logger.error(
"Amount too small to refund",
extra={
"amount_msats": amount_msats,
"amount_sats": amount_sats,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise ValueError("Amount too small to refund (less than 1 sat)")
# Atomically deduct the balance to avoid race conditions when multiple
# refunds are triggered concurrently.
stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.balance) >= amount_msats)
.values(balance=col(ApiKey.balance) - amount_msats)
)
result = await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
if result.rowcount == 0:
logger.error(
"Insufficient balance for refund",
extra={
"requested_msats": amount_msats,
"key_hash": key.hashed_key[:8] + "...",
"current_balance": key.balance,
},
)
raise ValueError("Insufficient balance.")
await session.refresh(key)
await delete_key_if_zero_balance(key, session)
if key.refund_address is None:
logger.error(
"Refund address not set", extra={"key_hash": key.hashed_key[:8] + "..."}
)
raise ValueError("Refund address not set.")
try:
result = await wallet().send_to_lnurl(key.refund_address, amount=amount_sats)
logger.info(
"Refund sent successfully",
extra={
"amount_sats": amount_sats,
"refund_address": key.refund_address[:20] + "..."
if len(key.refund_address) > 20
else key.refund_address,
"key_hash": key.hashed_key[:8] + "...",
"transaction_result": str(result),
},
)
return result
except Exception as e:
logger.error(
"Failed to send refund",
extra={
"error": str(e),
"error_type": type(e).__name__,
"amount_sats": amount_sats,
"refund_address": key.refund_address,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
async def x_cashu_refund(key: ApiKey, session: AsyncSession, unit: CurrencyUnit) -> str:
"""Process an X-Cashu refund token."""
logger.info(
"Processing X-Cashu refund",
extra={
"balance_msats": key.balance,
"unit": unit,
"key_hash": key.hashed_key[:8] + "...",
},
)
try:
refund_token = await wallet().send(key.balance, unit=unit)
logger.info(
"X-Cashu refund token created",
extra={
"amount": key.balance,
"unit": unit,
"key_hash": key.hashed_key[:8] + "...",
"token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
await session.delete(key)
await session.commit()
logger.info(
"X-Cashu refund completed", extra={"key_hash": key.hashed_key[:8] + "..."}
)
return refund_token
except Exception as e:
logger.error(
"Failed to create X-Cashu refund",
extra={
"error": str(e),
"error_type": type(e).__name__,
"balance": key.balance,
"unit": unit,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
async def redeem(cashu_token: str, lnurl: str) -> int:
"""Redeem a Cashu token and send to LNURL."""
logger.info(
"Starting token redemption for LNURL",
extra={
"token_preview": cashu_token[:20] + "..."
if len(cashu_token) > 20
else cashu_token,
"lnurl_preview": lnurl[:20] + "..." if len(lnurl) > 20 else lnurl,
},
)
try:
amount, unit = await wallet().redeem(cashu_token)
logger.info("Token redeemed for LNURL", extra={"amount": amount, "unit": unit})
unit = cast(CurrencyUnit, unit)
result = await wallet().send_to_lnurl(lnurl, amount=amount, unit=unit)
logger.info(
"Successfully sent to LNURL",
extra={
"amount": amount,
"unit": unit,
"lnurl_preview": lnurl[:20] + "..." if len(lnurl) > 20 else lnurl,
"transaction_result": str(result),
},
)
return amount
except Exception as e:
logger.error(
"Failed to redeem and send to LNURL",
extra={
"error": str(e),
"error_type": type(e).__name__,
"token_preview": cashu_token[:20] + "..."
if len(cashu_token) > 20
else cashu_token,
"lnurl_preview": lnurl[:20] + "..." if len(lnurl) > 20 else lnurl,
},
)
raise
+3
View File
@@ -0,0 +1,3 @@
from .logging import get_logger
__all__ = ["get_logger"]
+405
View File
@@ -0,0 +1,405 @@
import os
from datetime import datetime, timezone
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import HTMLResponse
from pydantic import BaseModel
from sqlmodel import select
from ..wallet import get_balance, send_token
from .db import ApiKey, create_session
admin_router = APIRouter(prefix="/admin", include_in_schema=False)
class WithdrawRequest(BaseModel):
amount: int
def login_form() -> str:
return """<!DOCTYPE html>
<html>
<head>
<style>
body {
font-family: Arial, sans-serif;
display: flex;
justify-content: center;
align-items: center;
height: 100vh;
margin: 0;
}
form {
display: flex;
flex-direction: column;
gap: 10px;
}
input[type="password"] {
padding: 8px;
}
button {
padding: 8px;
cursor: pointer;
}
</style>
<script>
function handleSubmit(e) {
e.preventDefault();
const password = document.getElementById('password').value;
document.cookie = `admin_password=${password}; path=/; max-age=86400`;
window.location.reload();
}
</script>
</head>
<body>
<form onsubmit="handleSubmit(event)">
<input type="password" id="password" placeholder="Admin Password" required>
<button type="submit">Login</button>
</form>
</body>
</html>
"""
def info(content: str) -> str:
return f"""<!DOCTYPE html>
<html>
<head>
<style>
body {{
font-family: Arial, sans-serif;
display: flex;
justify-content: center;
align-items: center;
height: 100vh;
margin: 0;
}}
</style>
</head>
<body>
<div style="text-align: center;">
{content}
</div>
</body>
</html>
"""
def admin_auth() -> str:
if os.getenv("ADMIN_PASSWORD", "") == "":
return info("Please set a secure ADMIN_PASSWORD= in your ENV variables.")
else:
return login_form()
async def dashboard(request: Request) -> str:
# fetch cashu / api-key data from database
async with create_session() as session:
result = await session.exec(select(ApiKey))
api_keys = result.all()
api_keys_table_rows = []
for key in api_keys:
expiry_time_utc = (
datetime.fromtimestamp(key.key_expiry_time, tz=timezone.utc)
if key.key_expiry_time is not None
else None
)
expiry_time_human_readable = (
expiry_time_utc.strftime("%Y-%m-%d %H:%M:%S") if expiry_time_utc else ""
)
api_keys_table_rows.append(
f"<tr><td>{key.hashed_key}</td><td>{key.balance}</td><td>{key.total_spent}</td><td>{key.total_requests}</td><td>{key.refund_address}</td><td>{'{} ({} UTC)'.format(key.key_expiry_time, expiry_time_human_readable) if key.key_expiry_time else key.key_expiry_time}</td></tr>"
)
# Calculate the total balance of all API keys using integer arithmetic to
# avoid rounding issues.
total_user_balance = sum(key.balance for key in api_keys) // 1000
# Fetch balance from cashu
current_balance = await get_balance("sat")
owner_balance = current_balance - total_user_balance
return f"""<!DOCTYPE html>
<html>
<head>
<style>
table {{
width: 100%;
border-collapse: collapse;
}}
th, td {{
border: 1px solid black;
padding: 8px;
text-align: left;
}}
button {{
padding: 8px 16px;
cursor: pointer;
background-color: #007bff;
color: white;
border: none;
border-radius: 4px;
margin-right: 10px;
}}
button:hover {{
background-color: #0056b3;
}}
button:disabled {{
background-color: #6c757d;
cursor: not-allowed;
}}
#token-result {{
margin-top: 20px;
padding: 15px;
background-color: #f8f9fa;
border: 1px solid #dee2e6;
border-radius: 4px;
word-break: break-all;
display: none;
max-width: 100%;
}}
#token-text {{
font-family: monospace;
font-size: 12px;
background-color: #e9ecef;
padding: 10px;
border-radius: 4px;
margin: 10px 0;
}}
.copy-btn {{
background-color: #28a745;
padding: 4px 8px;
font-size: 12px;
}}
.copy-btn:hover {{
background-color: #1e7e34;
}}
.refresh-btn {{
background-color: #ffc107;
color: black;
}}
.refresh-btn:hover {{
background-color: #e0a800;
}}
.modal {{
display: none;
position: fixed;
z-index: 1;
left: 0;
top: 0;
width: 100%;
height: 100%;
background-color: rgba(0,0,0,0.4);
}}
.modal-content {{
background-color: #fefefe;
margin: 15% auto;
padding: 20px;
border: 1px solid #888;
width: 300px;
border-radius: 8px;
text-align: center;
}}
.close {{
color: #aaa;
float: right;
font-size: 28px;
font-weight: bold;
cursor: pointer;
}}
.close:hover {{
color: black;
}}
input[type="number"] {{
width: 100%;
padding: 8px;
margin: 10px 0;
border: 1px solid #ddd;
border-radius: 4px;
}}
.warning {{
color: #dc3545;
font-weight: bold;
margin: 10px 0;
}}
</style>
<script>
function openWithdrawModal() {{
const modal = document.getElementById('withdraw-modal');
const amountInput = document.getElementById('withdraw-amount');
amountInput.value = {owner_balance};
modal.style.display = 'block';
}}
function closeWithdrawModal() {{
const modal = document.getElementById('withdraw-modal');
modal.style.display = 'none';
}}
function checkAmount() {{
const amount = parseInt(document.getElementById('withdraw-amount').value);
const warning = document.getElementById('withdraw-warning');
const ownerBalance = {owner_balance};
if (amount > ownerBalance && amount <= {current_balance}) {{
warning.style.display = 'block';
}} else {{
warning.style.display = 'none';
}}
}}
async function performWithdraw() {{
const amount = parseInt(document.getElementById('withdraw-amount').value);
const button = document.getElementById('confirm-withdraw-btn');
const tokenResult = document.getElementById('token-result');
if (!amount || amount <= 0) {{
alert('Please enter a valid amount');
return;
}}
if (amount > {current_balance}) {{
alert('Amount exceeds wallet balance');
return;
}}
button.disabled = true;
button.textContent = 'Withdrawing...';
try {{
const response = await fetch('/admin/withdraw', {{
method: 'POST',
headers: {{
'Content-Type': 'application/json',
}},
credentials: 'same-origin',
body: JSON.stringify({{ amount: amount }})
}});
if (response.ok) {{
const data = await response.json();
document.getElementById('token-text').textContent = data.token;
tokenResult.style.display = 'block';
closeWithdrawModal();
}} else {{
const errorData = await response.json();
alert('Failed to withdraw balance: ' + (errorData.detail || 'Unknown error'));
}}
}} catch (error) {{
alert('Error: ' + error.message);
}} finally {{
button.disabled = false;
button.textContent = 'Withdraw';
}}
}}
function copyToken() {{
const tokenText = document.getElementById('token-text');
navigator.clipboard.writeText(tokenText.textContent).then(() => {{
const copyBtn = document.getElementById('copy-btn');
const originalText = copyBtn.textContent;
copyBtn.textContent = 'Copied!';
setTimeout(() => {{
copyBtn.textContent = originalText;
}}, 2000);
}}).catch(err => {{
alert('Failed to copy token');
}});
}}
function refreshPage() {{
window.location.reload();
}}
window.onclick = function(event) {{
const modal = document.getElementById('withdraw-modal');
if (event.target == modal) {{
closeWithdrawModal();
}}
}}
</script>
</head>
<body>
<h1>Admin Dashboard</h1>
<h2>Current Cashu Balance</h2>
<p>Your Balance: {owner_balance} sats</p>
<p>The balance is calculated by subtracting the combined user balance from the total Cashu wallet balance.</p>
<p>Total Cashu Balance: {current_balance} sats</p>
<p>User Balance: {total_user_balance} sats</p>
<button id="withdraw-btn" onclick="openWithdrawModal()" {"disabled" if current_balance <= 0 else ""}>
Withdraw Balance
</button>
<button class="refresh-btn" onclick="refreshPage()">
Refresh Dashboard
</button>
<div id="withdraw-modal" class="modal">
<div class="modal-content">
<span class="close" onclick="closeWithdrawModal()">&times;</span>
<h3>Withdraw Balance</h3>
<p>Enter amount to withdraw (sats):</p>
<input type="number" id="withdraw-amount" min="1" max="{current_balance}" placeholder="Amount in sats" oninput="checkAmount()">
<p>Maximum: {current_balance} sats</p>
<p>Your recommended balance: {owner_balance} sats</p>
<div id="withdraw-warning" class="warning" style="display: none;">
⚠️ Warning: Withdrawing more than your balance will use user funds!
</div>
<button id="confirm-withdraw-btn" onclick="performWithdraw()">Withdraw</button>
<button onclick="closeWithdrawModal()" style="background-color: #6c757d;">Cancel</button>
</div>
</div>
<div id="token-result">
<strong>Withdrawal Token:</strong>
<div id="token-text"></div>
<button id="copy-btn" class="copy-btn" onclick="copyToken()">Copy Token</button>
<p><em>Save this token! It represents your withdrawn balance.</em></p>
</div>
<h2>User's API Keys</h2>
<table>
<tr>
<th>Hashed Key</th>
<th>Balance (mSats)</th>
<th>Total Spent (mSats)</th>
<th>Total Requests</th>
<th>Refund Address</th>
<th>Refund Time</th>
</tr>
{"".join(api_keys_table_rows)}
</table>
</body>
</html>
"""
@admin_router.get("/", response_class=HTMLResponse)
async def admin(request: Request) -> str:
admin_cookie = request.cookies.get("admin_password")
if admin_cookie and admin_cookie == os.getenv("ADMIN_PASSWORD"):
return await dashboard(request)
return admin_auth()
@admin_router.post("/withdraw")
async def withdraw(
request: Request, withdraw_request: WithdrawRequest
) -> dict[str, str]:
admin_cookie = request.cookies.get("admin_password")
if not admin_cookie or admin_cookie != os.getenv("ADMIN_PASSWORD"):
raise HTTPException(status_code=403, detail="Unauthorized")
current_balance = await get_balance("sat")
if withdraw_request.amount <= 0:
raise HTTPException(
status_code=400, detail="Withdrawal amount must be positive"
)
if withdraw_request.amount > current_balance:
raise HTTPException(status_code=400, detail="Insufficient wallet balance")
token = await send_token(withdraw_request.amount, "sat")
return {"token": token}
View File
@@ -8,6 +8,21 @@ from pathlib import Path
from typing import Any
from pythonjsonlogger import jsonlogger
from rich.logging import RichHandler
# Define custom TRACE level
TRACE_LEVEL = 5
logging.addLevelName(TRACE_LEVEL, "TRACE")
def trace(self: logging.Logger, message: str, *args: Any, **kwargs: Any) -> None:
"""Log with TRACE level"""
if self.isEnabledFor(TRACE_LEVEL):
self._log(TRACE_LEVEL, message, args, **kwargs)
# Add the trace method to Logger class
setattr(logging.Logger, "trace", trace)
class DailyRotatingFileHandler(logging.handlers.TimedRotatingFileHandler):
@@ -149,13 +164,33 @@ class SecurityFilter(logging.Filter):
def get_log_level() -> str:
"""Get log level from environment variable."""
return os.environ.get("LOG_LEVEL", "INFO").upper()
level = os.environ.get("LOG_LEVEL", "INFO").upper()
# Validate log level - if invalid, default to INFO
valid_levels = {"TRACE", "DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"}
if level not in valid_levels:
level = "INFO"
return level
def should_enable_console_logging() -> bool:
"""Check if console logging should be enabled."""
return os.environ.get("ENABLE_CONSOLE_LOGGING", "true").lower() in (
"true",
"1",
"yes",
)
def setup_logging() -> None:
"""Configure centralized logging for the application."""
log_level = get_log_level()
console_enabled = should_enable_console_logging()
# Determine which handlers to use
handlers = ["file"]
if console_enabled:
handlers.append("console")
LOGGING_CONFIG = {
"version": 1,
@@ -166,10 +201,6 @@ def setup_logging() -> None:
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s",
"datefmt": "%Y-%m-%d %H:%M:%S",
},
"standard": {
"format": "%(asctime)s [%(levelname)s] %(name)s v%(version)s: %(message)s",
"datefmt": "%Y-%m-%d %H:%M:%S",
},
},
"filters": {
"version_filter": {"()": VersionFilter},
@@ -177,13 +208,13 @@ def setup_logging() -> None:
},
"handlers": {
"console": {
"class": "logging.StreamHandler",
"()": RichHandler,
"level": log_level,
"formatter": "json"
if os.environ.get("LOG_FORMAT", "json").lower() == "json"
else "standard",
"stream": "ext://sys.stdout",
"filters": ["version_filter", "security_filter"],
"show_time": False,
"show_path": False,
"rich_tracebacks": True,
"markup": True,
"filters": ["security_filter"],
},
"file": {
"()": DailyRotatingFileHandler,
@@ -200,47 +231,56 @@ def setup_logging() -> None:
"loggers": {
"router": {
"level": log_level,
"handlers": ["console", "file"],
"handlers": handlers,
"propagate": False,
},
"router.payment": {
"level": log_level,
"handlers": ["console", "file"],
"handlers": handlers,
"propagate": False,
},
"router.cashu": {
"level": log_level,
"handlers": ["console", "file"],
"handlers": handlers,
"propagate": False,
},
"router.proxy": {
"level": log_level,
"handlers": ["console", "file"],
"handlers": handlers,
"propagate": False,
},
"router.auth": {
"level": log_level,
"handlers": ["console", "file"],
"handlers": handlers,
"propagate": False,
},
# Suppress verbose third-party logging
"httpx": {
"level": "WARNING",
"handlers": ["console"],
"handlers": ["console"] if console_enabled else [],
"propagate": False,
},
"httpcore": {
"level": "WARNING",
"handlers": ["console"],
"handlers": ["console"] if console_enabled else [],
"propagate": False,
},
"uvicorn.access": {
"level": "WARNING",
"handlers": ["console"] if console_enabled else [],
"propagate": False,
},
"uvicorn.error": {
"level": "INFO",
"handlers": ["console"],
"propagate": False,
},
"watchfiles.main": {"level": "WARNING", "handlers": [], "propagate": False},
},
"root": {
"level": log_level,
"handlers": ["console"] if console_enabled else [],
},
"root": {"level": log_level, "handlers": ["console"]},
}
os.makedirs("logs", exist_ok=True)
+11 -28
View File
@@ -6,14 +6,14 @@ from typing import AsyncGenerator
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from .account import wallet_router
from ..balance import balance_router, deprecated_wallet_router
from ..discovery import providers_router
from ..payment.models import MODELS, models_router, update_sats_pricing
from ..proxy import proxy_router
from ..wallet import check_for_refunds, periodic_payout
from .admin import admin_router
from .cashu import check_for_refunds, init_wallet, periodic_payout
from .db import init_db
from .discovery import providers_router
from .logging.logging_config import get_logger, setup_logging
from .models import MODELS, models_router, update_sats_pricing
from .proxy import proxy_router
from .logging import get_logger, setup_logging
# Initialize logging first
setup_logging()
@@ -28,20 +28,11 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
try:
await init_db()
logger.info("Database initialized successfully")
await init_wallet()
logger.info("Wallet initialized successfully")
pricing_task = asyncio.create_task(update_sats_pricing())
refund_task = asyncio.create_task(check_for_refunds())
payout_task = asyncio.create_task(periodic_payout())
logger.info(
"Background tasks started successfully",
extra={"tasks": ["pricing", "refunds", "payouts"]},
)
yield
except Exception as e:
@@ -86,13 +77,9 @@ app.add_middleware(
allow_headers=["*"],
)
logger.info(
"CORS middleware configured",
extra={"allowed_origins": os.environ.get("CORS_ORIGINS", "*").split(",")},
)
@app.get("/")
@app.get("/", include_in_schema=False)
@app.get("/v1/info")
async def info() -> dict:
logger.info("Info endpoint accessed")
return {
@@ -100,7 +87,7 @@ async def info() -> dict:
"description": app.description,
"version": __version__,
"npub": os.environ.get("NPUB", ""),
"mint": os.environ.get("MINT", ""),
"mints": os.environ.get("CASHU_MINTS", "").split(","),
"http_url": os.environ.get("HTTP_URL", ""),
"onion_url": os.environ.get("ONION_URL", ""),
"models": MODELS,
@@ -109,11 +96,7 @@ async def info() -> dict:
app.include_router(models_router)
app.include_router(admin_router)
app.include_router(wallet_router)
app.include_router(balance_router)
app.include_router(deprecated_wallet_router)
app.include_router(providers_router)
app.include_router(proxy_router)
logger.info(
"Application initialized successfully",
extra={"version": __version__, "routers_count": 5},
)
+163 -102
View File
@@ -2,8 +2,8 @@ import asyncio
import json
import os
import random
import re
import string
from typing import Any
import httpx
import websockets
@@ -17,59 +17,34 @@ def generate_subscription_id() -> str:
return "".join(random.choices(string.ascii_lowercase + string.digits, k=10))
def extract_onion_urls(content: str) -> list[str]:
"""Extract onion URLs from content."""
pattern = r"http?://[a-zA-Z0-9\-._~]+\.onion"
return re.findall(pattern, content)
async def query_nostr_relay_with_search(
search_term: str,
async def query_nostr_relay_for_providers(
relay_url: str,
kinds: list[int] | None = None,
pubkey: str | None = None,
limit: int = 1000,
timeout: int = 30,
) -> list[dict]:
) -> list[dict[str, Any]]:
"""
Query a Nostr relay and filter for events containing a search term.
Query a Nostr relay for provider announcements using RIP-02 spec.
Searches for kind 31338 events (Routstr Provider Announcements).
"""
if kinds is None:
kinds = [1]
events = []
# If searching for an npub mention, try tag-based search first
if search_term.startswith("nostr:npub"):
# Extract the npub and convert to hex
npub = search_term.replace("nostr:", "")
try:
# Convert npub to hex (you might need to implement or import this)
# For now, try tag-based search with the npub
filter_obj = {
"kinds": kinds,
"limit": limit,
"#p": [npub], # Posts that tag this pubkey
}
except Exception:
# If conversion fails, try regular search
filter_obj = {
"kinds": kinds,
"limit": limit,
}
else:
# Try relay's search functionality (NIP-50)
filter_obj = {
"kinds": kinds,
"search": search_term,
"limit": limit,
}
# Build filter according to RIP-02 spec
filter_obj: dict[str, Any] = {
"kinds": [31338], # RIP-02 Provider Announcement events
"limit": limit,
}
# If specific pubkey provided, filter by author
if pubkey:
filter_obj["authors"] = [pubkey]
sub_id = generate_subscription_id()
req_message = json.dumps(["REQ", sub_id, filter_obj])
try:
async with websockets.connect(relay_url, timeout=timeout) as websocket:
print(f"Connected to relay, sending request with filter: {filter_obj}")
print("Connected to relay, searching for kind 31338 events")
await websocket.send(req_message)
while True:
@@ -78,25 +53,14 @@ async def query_nostr_relay_with_search(
data = json.loads(message)
if data[0] == "EVENT" and data[1] == sub_id:
# For tag-based search, also check content
if search_term.startswith("nostr:npub"):
if search_term.lower() in data[2]["content"].lower():
print(f"Found matching event: {data[2]['id']}")
events.append(data[2])
else:
print(f"Found matching event: {data[2]['id']}")
events.append(data[2])
event = data[2]
print(f"Found provider announcement: {event['id']}")
events.append(event)
elif data[0] == "EOSE" and data[1] == sub_id:
print("Received EOSE message")
break
elif data[0] == "NOTICE":
print(f"Relay notice: {data[1]}")
# If search not supported, could break and try different approach
if "unrecognised filter item" in data[1] and "search" in str(
filter_obj
):
print("Search not supported on this relay")
break
except asyncio.TimeoutError:
print("Timeout waiting for message")
@@ -110,61 +74,159 @@ async def query_nostr_relay_with_search(
except Exception as e:
print(f"Query failed: {e}")
print(f"Query complete. Found {len(events)} matching events")
print(f"Query complete. Found {len(events)} provider announcements")
return events
async def get_cache() -> list[dict]:
def parse_provider_announcement(event: dict[str, Any]) -> dict[str, Any] | None:
"""
Parse a kind 31338 provider announcement event according to RIP-02 spec.
Returns structured provider data or None if invalid.
"""
try:
# Extract required tags according to RIP-02
tags = event.get("tags", [])
# Find required tags
endpoint_url = None
provider_name = None
d_tag = None
for tag in tags:
if len(tag) >= 2:
if tag[0] == "endpoint":
endpoint_url = tag[1]
elif tag[0] == "name":
provider_name = tag[1]
elif tag[0] == "d":
d_tag = tag[1]
# Validate required fields
if not endpoint_url or not provider_name or not d_tag:
print(
f"Invalid provider announcement - missing required tags: {event['id']}"
)
return None
# Extract optional tags
description = None
contact = None
pricing_url = None
supported_models = []
for tag in tags:
if len(tag) >= 2:
if tag[0] == "description":
description = tag[1]
elif tag[0] == "contact":
contact = tag[1]
elif tag[0] == "pricing":
pricing_url = tag[1]
elif tag[0] == "model":
supported_models.append(tag[1])
return {
"id": event["id"],
"pubkey": event["pubkey"],
"created_at": event["created_at"],
"d_tag": d_tag,
"endpoint_url": endpoint_url,
"name": provider_name,
"description": description,
"contact": contact,
"pricing_url": pricing_url,
"supported_models": supported_models,
"content": event.get("content", ""),
}
except Exception as e:
print(f"Error parsing provider announcement {event.get('id', 'unknown')}: {e}")
return None
async def get_cache() -> list[dict[str, Any]]:
return [] # TODO: Implement cache
async def fetch_onion(provider: str) -> dict:
"""Check if an onion service is healthy by making a GET request to its root."""
async def fetch_provider_health(endpoint_url: str) -> dict[str, Any]:
"""Check if a provider endpoint is healthy by making a GET request."""
try:
# Get Tor proxy URL from environment variable, default to local Tor SOCKS5 proxy
tor_proxy = os.getenv("TOR_PROXY_URL", "socks5://127.0.0.1:9050")
# Determine if we need Tor proxy based on .onion domain
is_onion = ".onion" in endpoint_url
# Set up client arguments conditionally
proxies = None
if is_onion:
# Get Tor proxy URL from environment variable
tor_proxy = os.getenv("TOR_PROXY_URL", "socks5://127.0.0.1:9050")
proxies = {"http://": tor_proxy, "https://": tor_proxy} # type: ignore[assignment]
# Configure httpx to use Tor SOCKS5 proxy
async with httpx.AsyncClient(
proxies={"http://": tor_proxy, "https://": tor_proxy}, # type: ignore
timeout=httpx.Timeout(30.0),
follow_redirects=True,
proxies=proxies, # type: ignore[arg-type]
) as client:
response = await client.get(provider)
# Consider 2xx and 3xx status codes as healthy
return {"status_code": response.status_code, "json": response.json()}
except Exception:
# Any exception means the service is not healthy
return {"status_code": 500, "json": {"error": "Failed to fetch onion"}}
# Try to fetch models endpoint first (common for AI providers)
models_url = f"{endpoint_url.rstrip('/')}/v1/models"
try:
response = await client.get(models_url)
if response.status_code == 200:
return {
"status_code": response.status_code,
"endpoint": "models",
"json": response.json(),
}
except Exception:
pass
# Fallback to root endpoint
response = await client.get(endpoint_url)
return {
"status_code": response.status_code,
"endpoint": "root",
"json": response.json()
if response.headers.get("content-type", "").startswith(
"application/json"
)
else {"message": "OK"},
}
except Exception as e:
return {
"status_code": 500,
"endpoint": "error",
"json": {"error": f"Failed to fetch provider: {str(e)}"},
}
@providers_router.get("/")
async def get_providers(include_json: bool = False) -> dict[str, list[dict | str]]:
npub = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s"
async def get_providers(
include_json: bool = False, pubkey: str | None = None
) -> dict[str, list[dict[str, Any]]]:
"""
Discover Routstr providers using RIP-02 specification.
Searches for kind 31338 provider announcement events on Nostr relays.
# Relays that support NIP-50 text search
search_relays = [
"wss://relay.nostr.band", # Known to support search
"wss://nostr.wine", # Known to support search
Reference: https://github.com/Routstr/protocol/blob/main/RIP-02.md
"""
# Default relays for provider discovery
discovery_relays = [
"wss://relay.nostr.band",
"wss://relay.damus.io",
"wss://nos.lol",
"wss://relay.routstr.com",
]
# Search for the mention format that appears in posts
search_term = f"nostr:{npub}"
all_events = []
event_ids = set() # To avoid duplicates
# Try multiple relays
for relay_url in search_relays:
print(f"\nTrying relay: {relay_url}")
# Query multiple relays for provider announcements
for relay_url in discovery_relays:
print(f"\nQuerying relay for providers: {relay_url}")
try:
events = await query_nostr_relay_with_search(
search_term=search_term,
events = await query_nostr_relay_for_providers(
relay_url=relay_url,
kinds=[1], # Text notes
limit=500,
pubkey=pubkey,
limit=100,
)
# Add unique events
@@ -173,35 +235,34 @@ async def get_providers(include_json: bool = False) -> dict[str, list[dict | str
event_ids.add(event["id"])
all_events.append(event)
print(f"Got {len(events)} events from {relay_url}")
# If we have enough events, we can stop
if len(all_events) >= 100:
break
print(f"Got {len(events)} provider announcements from {relay_url}")
except Exception as e:
print(f"Failed to query {relay_url}: {e}")
continue
print(f"Found {len(all_events)} total unique events mentioning routstr")
print(f"Found {len(all_events)} total unique provider announcements")
# Parse provider announcements according to RIP-02
providers = []
for event in all_events:
onion_urls = extract_onion_urls(event["content"])
providers.extend(onion_urls)
parsed_provider = parse_provider_announcement(event)
if parsed_provider:
providers.append(parsed_provider)
unique_providers = list(set(providers))
print(f"Parsed {len(providers)} valid provider announcements")
print(f"Found {len(unique_providers)} unique onion URLs")
print(unique_providers)
healthy_providers: list[dict | str] = []
for provider in unique_providers:
response = await fetch_onion(provider)
# Check provider health if requested
healthy_providers: list[dict[str, Any]] = []
for provider in providers:
endpoint_url = provider["endpoint_url"]
if include_json:
healthy_providers.append({provider: response["json"]})
health_check = await fetch_provider_health(endpoint_url)
provider_data = {"provider": provider, "health": health_check}
healthy_providers.append(provider_data)
else:
# Just return the provider info without health check
healthy_providers.append(provider)
return {"providers": healthy_providers}
View File
+8
View File
@@ -0,0 +1,8 @@
from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost
__all__ = [
"CostData",
"CostDataError",
"MaxCostData",
"calculate_cost",
]
+8 -7
View File
@@ -1,9 +1,10 @@
import math
import os
from pydantic import BaseModel
from router.logging.logging_config import get_logger
from router.models import MODELS
from ..core import get_logger
from .models import MODELS
logger = get_logger(__name__)
@@ -145,9 +146,9 @@ def calculate_cost(
input_tokens = response_data.get("usage", {}).get("prompt_tokens", 0)
output_tokens = response_data.get("usage", {}).get("completion_tokens", 0)
input_msats = int(round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 0))
output_msats = int(round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 0))
token_based_cost = int(round(input_msats + output_msats, 0))
input_msats = round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 3)
output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3)
token_based_cost = math.ceil(input_msats + output_msats)
logger.info(
"Calculated token-based cost",
@@ -163,7 +164,7 @@ def calculate_cost(
return CostData(
base_msats=0,
input_msats=input_msats,
output_msats=output_msats,
input_msats=int(input_msats),
output_msats=int(output_msats),
total_msats=token_based_cost,
)
+18 -176
View File
@@ -1,14 +1,12 @@
import base64
import json
import os
import cbor2
from fastapi import HTTPException, Response
from sixty_nuts.types import CurrencyUnit
from router.logging.logging_config import get_logger
from router.models import MODELS
from router.payment.cost_caculation import COST_PER_REQUEST, MODEL_BASED_PRICING
from ..core import get_logger
from ..wallet import deserialize_token_from_string
from .cost_caculation import COST_PER_REQUEST, MODEL_BASED_PRICING
from .models import MODELS
logger = get_logger(__name__)
@@ -49,17 +47,7 @@ def get_cost_per_request(model: str | None = None) -> int:
return COST_PER_REQUEST
def check_token_balance(headers: dict, body: dict) -> CurrencyUnit:
"""Check if the provided token has sufficient balance."""
logger.debug(
"Checking token balance",
extra={
"has_x_cashu": "x-cashu" in headers,
"has_authorization": "authorization" in headers,
"model": body.get("model", "unknown"),
},
)
def check_token_balance(headers: dict, body: dict) -> None:
if x_cashu := headers.get("x-cashu", None):
cashu_token = x_cashu
logger.debug(
@@ -100,173 +88,27 @@ def check_token_balance(headers: dict, body: dict) -> CurrencyUnit:
# Handle regular API keys (sk-*)
if cashu_token.startswith("sk-"):
logger.debug(
"Regular API key detected", extra={"key_preview": cashu_token[:10] + "..."}
)
return "sat"
return
cost = get_cost_per_request(model=body.get("model", None))
if cashu_token.startswith("cashuA"):
logger.debug("Processing CashuA token", extra={"required_cost_msats": cost})
token_obj = deserialize_token_from_string(cashu_token)
try:
_token = base64_token_json(cashu_token)
amount = sum(p["amount"] for t in _token["token"] for p in t["proofs"])
unit: CurrencyUnit = _token["unit"]
amount_msat = (
token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000
)
if unit == "sat":
amount *= 1000
logger.info(
"CashuA token parsed successfully",
extra={
"amount": amount,
"unit": _token["unit"],
"amount_msats": amount,
"required_cost_msats": cost,
"sufficient_balance": amount >= cost,
},
)
if amount < cost:
logger.warning(
"Insufficient token balance",
extra={
"amount_msats": amount,
"required_msats": cost,
"shortfall_msats": cost - amount,
"unit": unit,
},
)
raise HTTPException(status_code=413, detail="Insufficient balance")
except Exception as e:
logger.error(
"Failed to parse CashuA token",
extra={
"error": str(e),
"error_type": type(e).__name__,
"token_preview": cashu_token[:20] + "...",
},
)
raise HTTPException(status_code=401, detail="Invalid token format")
elif cashu_token.startswith("cashuB"):
logger.debug("Processing CashuB token", extra={"required_cost_msats": cost})
try:
_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
logger.info(
"CashuB token parsed successfully",
extra={
"amount": amount,
"unit": unit,
"amount_msats": amount,
"required_cost_msats": cost,
"sufficient_balance": amount >= cost,
},
)
if amount < cost:
logger.warning(
"Insufficient token balance",
extra={
"amount_msats": amount,
"required_msats": cost,
"shortfall_msats": cost - amount,
"unit": unit,
},
)
raise HTTPException(status_code=413, detail="Insufficient balance")
except Exception as e:
logger.error(
"Failed to parse CashuB token",
extra={
"error": str(e),
"error_type": type(e).__name__,
"token_preview": cashu_token[:20] + "...",
},
)
raise HTTPException(status_code=401, detail="Invalid token format")
else:
logger.error(
"Unknown token format",
extra={"token_prefix": cashu_token[:10] if cashu_token else "empty"},
)
raise HTTPException(status_code=401, detail="Unauthorized")
return unit
def base64_token_json(cashu_token: str) -> dict:
"""Decode a CashuA (JSON) token."""
logger.debug("Decoding CashuA token", extra={"token_length": len(cashu_token)})
try:
# 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)
logger.debug(
"CashuA token decoded successfully",
extra={
"token_proofs_count": sum(
len(t.get("proofs", [])) for t in token_data.get("token", [])
),
"unit": token_data.get("unit", "unknown"),
if cost > amount_msat:
raise HTTPException(
status_code=413,
detail={
"reason": "Insufficient balance",
"amount_required_msat": cost,
"model": body.get("model", "unknown"),
"type": "minimum_balance_required",
},
)
return token_data
except Exception as e:
logger.error(
"Failed to decode CashuA token",
extra={"error": str(e), "error_type": type(e).__name__},
)
raise
def base64_token_cbor(cashu_token: str) -> dict:
"""Decode a CashuB (CBOR) token."""
logger.debug("Decoding CashuB token", extra={"token_length": len(cashu_token)})
try:
encoded = cashu_token[6:] # Remove "cashuB"
encoded += "=" * ((-len(encoded)) % 4)
decoded_bytes = base64.urlsafe_b64decode(encoded)
token_data = cbor2.loads(decoded_bytes)
logger.debug(
"CashuB token decoded successfully",
extra={
"token_proofs_count": sum(
len(t.get("p", [])) for t in token_data.get("t", [])
),
"unit": token_data.get("u", "unknown"),
},
)
return token_data
except Exception as e:
logger.error(
"Failed to decode CashuB token",
extra={"error": str(e), "error_type": type(e).__name__},
)
raise
def get_max_cost_for_model(model: str) -> int:
"""Get the maximum cost for a specific model."""
@@ -147,7 +147,6 @@ async def update_sats_pricing() -> None:
break
@models_router.get("/models")
@models_router.get("/v1/models")
async def models() -> dict:
return {"data": MODELS}
+5 -4
View File
@@ -3,14 +3,15 @@ import os
import httpx
from .logging.logging_config import get_logger
from ..core import get_logger
logger = get_logger(__name__)
# artifical spread to cover conversion fees
EXCHANGE_FEE = float(os.environ.get("EXCHANGE_FEE", "1.005")) # 0.5% default
logger.info("Price module initialized", extra={"exchange_fee": EXCHANGE_FEE})
UPSTREAM_PROVIDER_FEE = float(
os.environ.get("UPSTREAM_PROVIDER_FEE", "1.05")
) # 5% default (e.g. openrouter charges 5% margin)
async def kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
@@ -98,7 +99,7 @@ async def btc_usd_ask_price() -> float:
raise ValueError("Unable to fetch BTC price from any exchange")
max_price = max(valid_prices)
final_price = max_price * EXCHANGE_FEE
final_price = max_price * EXCHANGE_FEE * UPSTREAM_PROVIDER_FEE
return final_price
+45 -58
View File
@@ -1,21 +1,20 @@
import json
import traceback
from typing import AsyncGenerator, Literal, cast
from typing import AsyncGenerator
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.logging.logging_config import get_logger
from router.payment.cost_caculation import (
from ..core import get_logger
from ..wallet import CurrencyUnit, recieve_token, send_token
from .cost_caculation import (
CostData,
CostDataError,
MaxCostData,
calculate_cost,
)
from router.payment.helpers import (
from .helpers import (
UPSTREAM_BASE_URL,
create_error_response,
get_max_cost_for_model,
@@ -42,26 +41,47 @@ async def x_cashu_handler(
try:
headers = dict(request.headers)
amount, unit = await redeem_token(x_cashu_token)
amount, unit, mint = await recieve_token(x_cashu_token)
headers = prepare_upstream_headers(dict(request.headers))
logger.info(
"X-Cashu token redeemed successfully",
extra={"amount": amount, "unit": unit, "path": path},
extra={"amount": amount, "unit": unit, "path": path, "mint": mint},
)
return await forward_to_upstream(request, path, headers, amount, unit)
except Exception as e:
error_message = str(e)
logger.error(
"X-Cashu payment request failed",
extra={
"error": str(e),
"error": error_message,
"error_type": type(e).__name__,
"path": path,
"method": request.method,
},
)
raise
# Handle specific CASHU errors with appropriate HTTP status codes
if "already spent" in error_message.lower():
return create_error_response(
"token_already_spent",
"The provided CASHU token has already been spent",
400,
)
elif "invalid token" in error_message.lower():
return create_error_response(
"invalid_token", "The provided CASHU token is invalid", 400
)
elif "mint error" in error_message.lower():
return create_error_response(
"mint_error", f"CASHU mint error: {error_message}", 422
)
else:
# Generic error for other cases
return create_error_response(
"cashu_error", f"CASHU token processing failed: {error_message}", 400
)
async def forward_to_upstream(
@@ -303,7 +323,13 @@ async def handle_streaming_response(
try:
cost_data = await get_cost(response_data)
if cost_data:
refund_amount = amount - cost_data.total_msats
if unit == "msat":
refund_amount = amount - cost_data.total_msats
elif unit == "sat":
refund_amount = amount - (cost_data.total_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
if refund_amount > 0:
logger.info(
"Processing refund for streaming response",
@@ -405,7 +431,12 @@ async def handle_non_streaming_response(
if "content-encoding" in response_headers:
del response_headers["content-encoding"]
refund_amount = amount - cost_data.total_msats
if unit == "msat":
refund_amount = amount - cost_data.total_msats
elif unit == "sat":
refund_amount = amount - (cost_data.total_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
logger.info(
"Processing non-streaming response cost calculation",
@@ -454,7 +485,7 @@ async def handle_non_streaming_response(
# Emergency refund with small deduction for processing
emergency_refund = amount
refund_token = await wallet().send(emergency_refund)
refund_token = await send_token(emergency_refund, unit=unit)
response.headers["X-Cashu"] = refund_token
logger.warning(
@@ -528,50 +559,6 @@ async def get_cost(response_data: dict) -> MaxCostData | CostData | None:
)
async def redeem_token(x_cashu_token: str) -> tuple[int, Literal["sat", "msat"]]:
"""Redeem X-Cashu token and return amount and unit."""
logger.debug(
"Redeeming X-Cashu token",
extra={
"token_preview": x_cashu_token[:20] + "..."
if len(x_cashu_token) > 20
else x_cashu_token
},
)
try:
result = await wallet().redeem(x_cashu_token)
amount, unit = cast(tuple[int, Literal["sat", "msat"]], result)
logger.info(
"X-Cashu token redeemed successfully",
extra={"amount": amount, "unit": unit},
)
return amount, unit
except Exception as e:
logger.error(
"X-Cashu token redemption failed",
extra={
"error": str(e),
"error_type": type(e).__name__,
"token_preview": x_cashu_token[:20] + "..."
if len(x_cashu_token) > 20
else x_cashu_token,
},
)
raise HTTPException(
status_code=401,
detail={
"error": {
"message": f"Invalid or expired Cashu key: {str(e)}",
"type": "invalid_request_error",
"code": "invalid_api_key",
}
},
)
async def send_refund(amount: int, unit: CurrencyUnit, mint: str | None = None) -> str:
"""Send a refund using Cashu tokens."""
logger.debug(
@@ -583,7 +570,7 @@ async def send_refund(amount: int, unit: CurrencyUnit, mint: str | None = None)
for attempt in range(max_retries):
try:
refund_token = await wallet().send(amount, unit=unit, mint_url=mint)
refund_token = await send_token(amount, unit=unit, mint_url=mint)
logger.info(
"Refund token created successfully",
+10 -22
View File
@@ -7,22 +7,21 @@ import httpx
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request
from fastapi.responses import Response, StreamingResponse
from router.logging.logging_config import get_logger
from router.payment.helpers import (
UPSTREAM_BASE_URL,
check_token_balance,
create_error_response,
prepare_upstream_headers,
)
from router.payment.x_cashu import x_cashu_handler
from .auth import (
adjust_payment_for_tokens,
pay_for_request,
revert_pay_for_request,
validate_bearer_key,
)
from .db import ApiKey, AsyncSession, create_session, get_session
from .core import get_logger
from .core.db import ApiKey, AsyncSession, create_session, get_session
from .payment.helpers import (
UPSTREAM_BASE_URL,
check_token_balance,
create_error_response,
prepare_upstream_headers,
)
from .payment.x_cashu import x_cashu_handler
logger = get_logger(__name__)
proxy_router = APIRouter()
@@ -480,18 +479,7 @@ async def proxy(
media_type="application/json",
)
# Check token balance for all requests to get currency unit
try:
unit = check_token_balance(headers, request_body_dict)
logger.debug(
"Token balance check completed", extra={"path": path, "unit": unit}
)
except HTTPException as e:
logger.warning(
"Token balance check failed",
extra={"path": path, "status_code": e.status_code, "detail": str(e.detail)},
)
raise
check_token_balance(headers, request_body_dict)
# Handle authentication
if x_cashu := headers.get("x-cashu", None):
+164
View File
@@ -0,0 +1,164 @@
import os
from typing import Literal
from cashu.core.base import Token
from cashu.wallet.helpers import deserialize_token_from_string, send
from cashu.wallet.wallet import Wallet
from .core import db, get_logger
logger = get_logger(__name__)
CurrencyUnit = Literal["sat", "msat"]
TRUSTED_MINTS = os.environ["CASHU_MINTS"].split(",")
PRIMARY_MINT_URL = TRUSTED_MINTS[0]
async def get_balance(unit: CurrencyUnit) -> int:
wallet = await Wallet.with_db(
PRIMARY_MINT_URL,
db=".wallet",
load_all_keysets=True,
unit=unit,
)
await wallet.load_proofs()
return wallet.available_balance.amount
async def recieve_token(
token: str,
) -> tuple[int, CurrencyUnit, str]: # amount, unit, mint_url
token_obj = deserialize_token_from_string(token)
if len(token_obj.keysets) > 1:
raise ValueError("Multiple keysets per token currently not supported")
wallet = await Wallet.with_db(
token_obj.mint,
db=".wallet",
load_all_keysets=True,
unit=token_obj.unit,
)
await wallet.load_mint(token_obj.keysets[0])
if token_obj.mint not in TRUSTED_MINTS:
return await swap_to_primary_mint(token_obj, wallet)
await wallet.redeem(token_obj.proofs)
return token_obj.amount, token_obj.unit, token_obj.mint
async def send_token(
amount: int, unit: CurrencyUnit, mint_url: str | None = None
) -> str:
wallet = await Wallet.with_db(
mint_url or PRIMARY_MINT_URL,
db=".wallet",
load_all_keysets=True,
unit=unit,
)
balance, token = await send(wallet, amount=amount, lock="", legacy=False)
return token
async def swap_to_primary_mint(
token_obj: Token, token_wallet: Wallet
) -> tuple[int, CurrencyUnit, str]:
logger.info(
"swap_to_primary_mint",
extra={
"mint": token_obj.mint,
"amount": token_obj.amount,
"unit": token_obj.unit,
},
)
if token_obj.unit == "sat":
amount_msat = token_obj.amount * 1000
elif token_obj.unit == "msat":
amount_msat = token_obj.amount
else:
raise ValueError("Invalid unit")
estimated_fee_sat = max(amount_msat // 1000 * 0.01, 2)
amount_msat_after_fee = amount_msat - estimated_fee_sat * 1000
primary_wallet = await Wallet.with_db(
PRIMARY_MINT_URL, db=".wallet", load_all_keysets=True, unit="sat"
)
await primary_wallet.load_mint()
minted_amount = amount_msat_after_fee // 1000
mint_quote = await primary_wallet.request_mint(minted_amount)
melt_quote = await token_wallet.melt_quote(mint_quote.request)
_ = await token_wallet.melt(
proofs=token_obj.proofs,
invoice=mint_quote.request,
fee_reserve_sat=melt_quote.fee_reserve,
quote_id=melt_quote.quote,
)
_ = await primary_wallet.mint(minted_amount, quote_id=mint_quote.quote)
return minted_amount, "sat", PRIMARY_MINT_URL
async def credit_balance(
cashu_token: str, key: db.ApiKey, session: db.AsyncSession
) -> int:
amount, unit, mint_url = await recieve_token(cashu_token)
if unit == "sat":
amount = amount * 1000
if mint_url != PRIMARY_MINT_URL:
raise ValueError("Mint URL is not supported by this proxy")
key.balance += amount
session.add(key)
await session.commit()
logger.info(
"Cashu token successfully redeemed and stored",
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
)
return amount
async def send_to_lnurl(amount: int, unit: CurrencyUnit, lnurl: str) -> dict[str, int]:
raise NotImplementedError
async def check_for_refunds() -> None:
logger.warning("check_for_refunds, temporary not implemented")
async def periodic_payout() -> None:
logger.warning("periodic_payout, temporary not implemented")
# class Proof:
# """
# Represents an ecash bill
# """
# def redeem_to_proofs(self, token: str) -> list[Proof]:
# raise NotImplementedError
# class Payment:
# """
# Stores all cashu payment related data
# """
# def __init__(self, token: str) -> None:
# self.initial_token = token
# amount, unit, mint_url = self.parse_token(token)
# self.amount = amount
# self.unit = unit
# self.mint_url = mint_url
# self.claimed_proofs = redeem_to_proofs(token)
# def parse_token(self, token: str) -> tuple[int, CurrencyUnit, str]:
# raise NotImplementedError
# def refund_full(self) -> None:
# raise NotImplementedError
# def refund_partial(self, amount: int) -> None:
# raise NotImplementedError
+268
View File
@@ -0,0 +1,268 @@
#!/usr/bin/env python3
"""
Simple Python function to publish one provider listing to a nostr relay
according to the RIP-02 specification.
Based on: https://github.com/Routstr/protocol/blob/main/RIP-02.md
Event Kind: 31338 (Routstr Provider Announcements)
"""
import asyncio
import hashlib
import json
import time
from typing import Any
import secp256k1
import websockets
def create_provider_announcement_event(
private_key_hex: str,
provider_name: str,
endpoint_url: str,
d_tag: str,
description: str | None = None,
contact: str | None = None,
pricing_url: str | None = None,
supported_models: list[str] | None = None,
) -> dict[str, Any]:
"""
Create a RIP-02 compliant provider announcement event.
Args:
private_key_hex: 32-byte hex private key for signing
provider_name: Human readable name for the provider
endpoint_url: Base URL for the provider's API endpoint
d_tag: Unique identifier for this provider (required for addressable events)
description: Optional description of the provider
contact: Optional contact information
pricing_url: Optional URL to pricing information
supported_models: Optional list of supported model names
Returns:
Complete signed nostr event ready for publishing
"""
# Convert hex private key to secp256k1 PrivateKey object
private_key = secp256k1.PrivateKey(bytes.fromhex(private_key_hex))
public_key = private_key.pubkey.serialize(compressed=True)[
1:
] # Remove 0x02/0x03 prefix
# Build required tags according to RIP-02
tags = [
["d", d_tag], # Required for addressable events (kind 30000-39999)
["endpoint", endpoint_url],
["name", provider_name],
]
# Add optional tags if provided
if description:
tags.append(["description", description])
if contact:
tags.append(["contact", contact])
if pricing_url:
tags.append(["pricing", pricing_url])
if supported_models:
for model in supported_models:
tags.append(["model", model])
# Create the event structure
created_at = int(time.time())
event_data = [
0, # Reserved field
public_key.hex(), # Public key as hex
created_at, # Unix timestamp
31338, # Kind for RIP-02 Provider Announcements
tags, # Tags array
"", # Content (empty for provider announcements)
]
# Serialize event data for hashing
event_json = json.dumps(event_data, separators=(",", ":"), ensure_ascii=False)
# Calculate event ID (SHA256 hash)
event_id = hashlib.sha256(event_json.encode("utf-8")).hexdigest()
# Sign the event ID
signature = private_key.ecdsa_sign(bytes.fromhex(event_id), raw=True)
signature_der = private_key.ecdsa_serialize(signature)
# Create the final event
event = {
"id": event_id,
"pubkey": public_key.hex(),
"created_at": created_at,
"kind": 31338,
"tags": tags,
"content": "",
"sig": signature_der.hex(),
}
return event
async def publish_provider_to_relay(
relay_url: str, event: dict[str, Any], timeout: int = 30
) -> bool:
"""
Publish a provider announcement event to a nostr relay.
Args:
relay_url: WebSocket URL of the nostr relay (e.g., "wss://relay.damus.io")
event: Complete signed nostr event to publish
timeout: Connection timeout in seconds
Returns:
True if successfully published, False otherwise
"""
try:
async with websockets.connect(relay_url, timeout=timeout) as websocket:
# Send EVENT message
event_message = json.dumps(["EVENT", event])
await websocket.send(event_message)
print(f"Published event {event['id']} to {relay_url}")
# Wait for OK response
try:
response = await asyncio.wait_for(websocket.recv(), timeout=5)
data = json.loads(response)
if data[0] == "OK" and data[1] == event["id"]:
if data[2]: # True means accepted
print(
f"✅ Event accepted by relay: {data[3] if len(data) > 3 else ''}"
)
return True
else:
print(
f"❌ Event rejected by relay: {data[3] if len(data) > 3 else ''}"
)
return False
elif data[0] == "NOTICE":
print(f"📢 Relay notice: {data[1]}")
return False
else:
print(f"🤔 Unexpected response: {data}")
return False
except asyncio.TimeoutError:
print("⏰ No response from relay within timeout")
return False
except Exception as e:
print(f"💥 Failed to publish to {relay_url}: {e}")
return False
async def publish_provider_listing(
private_key_hex: str,
provider_name: str,
endpoint_url: str,
d_tag: str,
relay_urls: list[str] | None = None,
description: str | None = None,
contact: str | None = None,
pricing_url: str | None = None,
supported_models: list[str] | None = None,
) -> dict[str, bool]:
"""
Complete function to create and publish a provider listing to nostr relays.
Args:
private_key_hex: 32-byte hex private key for signing
provider_name: Human readable name for the provider
endpoint_url: Base URL for the provider's API endpoint
d_tag: Unique identifier for this provider
relay_urls: List of relay URLs to publish to (uses defaults if None)
description: Optional description of the provider
contact: Optional contact information
pricing_url: Optional URL to pricing information
supported_models: Optional list of supported model names
Returns:
Dictionary mapping relay URLs to success status
"""
# Use default relays if none provided
if relay_urls is None:
relay_urls = [
"wss://relay.nostr.band",
"wss://relay.damus.io",
"wss://relay.routstr.com",
]
# Create the provider announcement event
event = create_provider_announcement_event(
private_key_hex=private_key_hex,
provider_name=provider_name,
endpoint_url=endpoint_url,
d_tag=d_tag,
description=description,
contact=contact,
pricing_url=pricing_url,
supported_models=supported_models,
)
print(f"📝 Created provider announcement event: {event['id']}")
print(f"🔑 Public key: {event['pubkey']}")
print(f"🏷️ Provider: {provider_name}")
print(f"🌐 Endpoint: {endpoint_url}")
print()
# Publish to all specified relays
results = {}
tasks = []
for relay_url in relay_urls:
task = publish_provider_to_relay(relay_url, event)
tasks.append((relay_url, task))
# Execute all publishing tasks concurrently
for relay_url, task in tasks:
try:
success = await task
results[relay_url] = success
except Exception as e:
print(f"💥 Failed to publish to {relay_url}: {e}")
results[relay_url] = False
return results
# Example usage
async def main() -> None:
"""Example of how to use the provider publishing function."""
# Example private key (DO NOT use this in production!)
private_key = "3185a47e3802f956ca207b46c8d6b8b5c5dbad53a5ca29816050e9b66badc33c"
# Example provider information
provider_name = "My AI Provider"
endpoint_url = "https://api.myaiprovider.com"
d_tag = "my-ai-provider-v1" # Unique identifier
description = "High-quality AI models with competitive pricing"
contact = "admin@myaiprovider.com"
pricing_url = "https://myaiprovider.com/pricing"
supported_models = ["gpt-4o", "claude-3-sonnet", "llama-3.1-70b"]
# Publish to relays
results = await publish_provider_listing(
private_key_hex=private_key,
provider_name=provider_name,
endpoint_url=endpoint_url,
d_tag=d_tag,
description=description,
contact=contact,
pricing_url=pricing_url,
supported_models=supported_models,
)
# Print results
print("\n📊 Publishing Results:")
for relay_url, success in results.items():
status = "✅ Success" if success else "❌ Failed"
print(f" {relay_url}: {status}")
if __name__ == "__main__":
asyncio.run(main())
+1 -2
View File
@@ -28,9 +28,8 @@ To run specific test files:
```bash
pytest tests/test_main.py
pytest tests/test_account.py
pytest tests/test_proxy.py
pytest tests/test_models.py
pytest tests/test_proxy.py
```
To run only async tests:
+16 -72
View File
@@ -1,7 +1,7 @@
import asyncio
import os
from typing import AsyncGenerator, Generator
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import patch
import pytest
import pytest_asyncio
@@ -14,14 +14,14 @@ from sqlmodel.ext.asyncio.session import AsyncSession
# Save original environment variables
ORIGINAL_ENV = os.environ.copy()
# Set test environment variables before importing the app
# Set test environment variables BEFORE importing the app
TEST_ENV = {
"UPSTREAM_BASE_URL": "https://api.example.com",
"UPSTREAM_API_KEY": "test-upstream-key",
"NAME": "TestRoutstrNode",
"DESCRIPTION": "Test Node",
"NPUB": "npub1test",
"MINT": "https://test.mint.com",
"CASHU_MINTS": "https://test.mint.com",
"HTTP_URL": "http://test.example.com",
"ONION_URL": "http://test.onion",
"CORS_ORIGINS": "*",
@@ -36,28 +36,9 @@ TEST_ENV = {
# Apply test environment
os.environ.update(TEST_ENV)
# Mock the Wallet class from sixty_nuts before importing the app
with patch("sixty_nuts.Wallet") as mock_wallet_class:
# Create a mock wallet instance
mock_wallet = AsyncMock()
mock_wallet.__aenter__ = AsyncMock(return_value=mock_wallet)
mock_wallet.__aexit__ = AsyncMock(return_value=None)
# Mock wallet state
mock_state = MagicMock()
mock_state.balance = 1000 # Balance in sats
mock_wallet.fetch_wallet_state = AsyncMock(return_value=mock_state)
# Mock other wallet methods
mock_wallet.send_to_lnurl = AsyncMock(return_value=100)
mock_wallet.redeem = AsyncMock(return_value=(1, "msat"))
mock_wallet.send = AsyncMock(return_value="cashu:token123")
# Make the Wallet class return our mock when instantiated
mock_wallet_class.return_value = mock_wallet
from router.db import get_session
from router.main import app
# Now import modules that depend on environment variables
from router.core.db import get_session # noqa: E402
from router.core.main import app # noqa: E402
@pytest.fixture(scope="session")
@@ -98,28 +79,9 @@ async def test_session(test_engine: AsyncEngine) -> AsyncGenerator[AsyncSession,
def test_client() -> Generator[TestClient, None, None]:
"""Create a test client for the FastAPI app."""
with patch.dict(os.environ, TEST_ENV, clear=True):
with patch("sixty_nuts.Wallet") as mock_wallet_class:
# Create a mock wallet instance
mock_wallet = AsyncMock()
mock_wallet.__aenter__ = AsyncMock(return_value=mock_wallet)
mock_wallet.__aexit__ = AsyncMock(return_value=None)
# Mock wallet state
mock_state = MagicMock()
mock_state.balance = 1000 # Balance in sats
mock_wallet.fetch_wallet_state = AsyncMock(return_value=mock_state)
# Mock other wallet methods
mock_wallet.send_to_lnurl = AsyncMock(return_value=100)
mock_wallet.redeem = AsyncMock(return_value=(1, "msat"))
mock_wallet.send = AsyncMock(return_value="cashu:token123")
# Make the Wallet class return our mock when instantiated
mock_wallet_class.return_value = mock_wallet
with patch("router.models.update_sats_pricing") as mock_update:
mock_update.return_value = None
yield TestClient(app)
with patch("router.payment.models.update_sats_pricing") as mock_update:
mock_update.return_value = None
yield TestClient(app)
@pytest_asyncio.fixture
@@ -133,32 +95,14 @@ async def async_client(test_session: AsyncSession) -> AsyncGenerator[AsyncClient
# Mock startup tasks
with patch.dict(os.environ, TEST_ENV, clear=True):
with patch("sixty_nuts.Wallet") as mock_wallet_class:
# Create a mock wallet instance
mock_wallet = AsyncMock()
mock_wallet.__aenter__ = AsyncMock(return_value=mock_wallet)
mock_wallet.__aexit__ = AsyncMock(return_value=None)
with patch("router.payment.models.update_sats_pricing") as mock_update:
mock_update.return_value = None
# Mock wallet state
mock_state = MagicMock()
mock_state.balance = 1000 # Balance in sats
mock_wallet.fetch_wallet_state = AsyncMock(return_value=mock_state)
# Mock other wallet methods
mock_wallet.send_to_lnurl = AsyncMock(return_value=100)
mock_wallet.redeem = AsyncMock(return_value=(1, "msat"))
mock_wallet.send = AsyncMock(return_value="cashuAoken123")
# Make the Wallet class return our mock when instantiated
mock_wallet_class.return_value = mock_wallet
with patch("router.models.update_sats_pricing") as mock_update:
mock_update.return_value = None
async with AsyncClient(
transport=ASGITransport(app=app), base_url="http://test"
) as client:
yield client
async with AsyncClient(
transport=ASGITransport(app=app), # type: ignore
base_url="http://test",
) as client:
yield client
app.dependency_overrides.clear()
+2 -2
View File
@@ -13,7 +13,7 @@ async def test_root_endpoint(async_client: AsyncClient) -> None:
"NAME": "TestRoutstrNode",
"DESCRIPTION": "Test Node",
"NPUB": "npub1test",
"MINT": "https://test.mint.com",
"CASHU_MINTS": "https://test.mint.com,https://test.mint2.com",
"HTTP_URL": "http://test.example.com",
"ONION_URL": "http://test.onion",
}
@@ -28,7 +28,7 @@ async def test_root_endpoint(async_client: AsyncClient) -> None:
assert "name" in data
assert "description" in data
assert "npub" in data
assert "mint" in data
assert "mints" in data
assert "http_url" in data
assert "onion_url" in data
+4 -4
View File
@@ -4,7 +4,7 @@ from unittest.mock import AsyncMock, patch
import pytest
from router.models import (
from router.payment.models import (
MODELS,
Architecture,
Model,
@@ -49,7 +49,7 @@ async def test_update_sats_pricing_calculation(sample_model: Model) -> None:
"""Test that sats pricing is calculated correctly."""
# Mock the sats_usd_ask_price function
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
"router.payment.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD
@@ -137,7 +137,7 @@ async def test_update_sats_pricing_without_top_provider() -> None:
)
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
"router.payment.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD
@@ -189,7 +189,7 @@ async def test_update_sats_pricing_without_top_provider() -> None:
async def test_update_sats_pricing_handles_errors() -> None:
"""Test that update_sats_pricing handles errors gracefully."""
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
"router.payment.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.side_effect = Exception("API Error")
-353
View File
@@ -1,353 +0,0 @@
import json
import os
import uuid
from typing import AsyncGenerator
from unittest.mock import AsyncMock, patch
import pytest
import pytest_asyncio
from httpx import AsyncClient
from router.db import ApiKey, AsyncSession
@pytest_asyncio.fixture
async def api_key_with_balance(test_session: AsyncSession) -> ApiKey:
"""Create an API key with sufficient balance."""
unique_id = str(uuid.uuid4())[:8]
key = ApiKey(
hashed_key=f"test-hashed-key-{unique_id}",
balance=10000000, # 10,000 sats in msats
refund_address=None,
total_spent=0,
total_requests=0,
)
test_session.add(key)
await test_session.commit()
await test_session.refresh(key)
return key
@pytest.mark.asyncio
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"
in response.json()["detail"]["error"]["message"]
)
@pytest.mark.asyncio
async def test_proxy_with_insufficient_balance(
async_client: AsyncClient, test_session: AsyncSession
) -> None:
"""Test proxy request with insufficient balance."""
# Create key with minimal balance
unique_id = str(uuid.uuid4())[:8]
key = ApiKey(
hashed_key=f"low-balance-key-{unique_id}",
balance=100, # Only 0.1 sats
refund_address=None,
total_spent=0,
total_requests=0,
)
test_session.add(key)
await test_session.commit()
# Mock the models.json check
with patch("os.path.exists", return_value=False):
response = await async_client.post(
"/v1/chat/completions",
headers={"Authorization": f"Bearer sk-{key.hashed_key}"},
json={"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]},
)
assert response.status_code == 402
assert "Insufficient balance" in response.json()["detail"]["error"]["message"]
@pytest.mark.asyncio
async def test_proxy_invalid_json_body(
async_client: AsyncClient, api_key_with_balance: ApiKey
) -> None:
"""Test proxy request with invalid JSON body."""
response = await async_client.post(
"/v1/chat/completions",
headers={
"Authorization": f"Bearer sk-{api_key_with_balance.hashed_key}",
"Content-Type": "application/json",
},
content=b'{"invalid": json",}', # Invalid JSON
)
assert response.status_code == 400
error_data = response.json()
assert error_data["error"]["type"] == "invalid_request_error"
assert error_data["error"]["code"] == "invalid_json"
@pytest.mark.asyncio
async def test_proxy_successful_request_mock(
async_client: AsyncClient, api_key_with_balance: ApiKey, test_session: AsyncSession
) -> None:
"""Test successful proxy request with mocked upstream."""
mock_response_data = {
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "gpt-4",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Hello! How can I help you?",
},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 9, "completion_tokens": 10, "total_tokens": 19},
}
with patch("httpx.AsyncClient") as mock_client_class:
mock_client = AsyncMock()
mock_client_class.return_value = mock_client
# Create a mock response
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.aread = AsyncMock(
return_value=json.dumps(mock_response_data).encode()
)
mock_response.aiter_bytes = AsyncMock()
mock_response.aclose = AsyncMock()
mock_client.send = AsyncMock(return_value=mock_response)
mock_client.build_request = AsyncMock()
mock_client.aclose = AsyncMock()
# Also mock the models.json check and pay_out
with patch("os.path.exists", return_value=False):
with patch("router.cashu.pay_out") as mock_payout:
mock_payout.return_value = None
response = await async_client.post(
"/v1/chat/completions",
headers={
"Authorization": f"Bearer sk-{api_key_with_balance.hashed_key}"
},
json={
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
},
)
assert response.status_code == 200
response_json = response.json()
# Verify the response includes the original data plus cost
assert response_json["id"] == "chatcmpl-123"
assert "cost" in response_json
assert response_json["cost"]["total_msats"] >= 0
# Verify balance was deducted
await test_session.refresh(api_key_with_balance)
assert api_key_with_balance.balance < 10000000
assert api_key_with_balance.total_requests == 1
@pytest.mark.asyncio
async def test_proxy_streaming_response(
async_client: AsyncClient, api_key_with_balance: ApiKey
) -> None:
"""Test proxy request with streaming response."""
# Mock SSE stream chunks
stream_chunks = [
b'data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"gpt-4","choices":[{"delta":{"content":"Hello"},"index":0}]}\n\n',
b'data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"gpt-4","choices":[{"delta":{"content":" there!"},"index":0}]}\n\n',
b'data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"gpt-4","choices":[{"delta":{},"index":0,"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":3,"total_tokens":12}}\n\n',
b"data: [DONE]\n\n",
]
async def mock_aiter_bytes() -> AsyncGenerator[bytes, None]:
for chunk in stream_chunks:
yield chunk
with patch("httpx.AsyncClient") as mock_client_class:
mock_client = AsyncMock()
mock_client_class.return_value = mock_client
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"}
mock_response.aiter_bytes = lambda: mock_aiter_bytes()
mock_response.aclose = AsyncMock()
mock_client.send = AsyncMock(return_value=mock_response)
mock_client.build_request = AsyncMock()
mock_client.aclose = AsyncMock()
with patch("os.path.exists", return_value=False):
with patch("router.cashu.pay_out") as mock_payout:
mock_payout.return_value = None
response = await async_client.post(
"/v1/chat/completions",
headers={
"Authorization": f"Bearer sk-{api_key_with_balance.hashed_key}"
},
json={
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"stream": True,
},
)
assert response.status_code == 200
assert response.headers["content-type"] == "text/event-stream"
@pytest.mark.asyncio
async def test_proxy_handles_upstream_errors(
async_client: AsyncClient, api_key_with_balance: ApiKey
) -> None:
"""Test proxy handles upstream connection errors gracefully."""
with patch("httpx.AsyncClient") as mock_client_class:
mock_client = AsyncMock()
mock_client_class.return_value = mock_client
# Simulate connection error
mock_client.send.side_effect = Exception("Connection refused")
mock_client.build_request = AsyncMock()
mock_client.aclose = AsyncMock()
with patch("os.path.exists", return_value=False):
response = await async_client.post(
"/v1/chat/completions",
headers={
"Authorization": f"Bearer sk-{api_key_with_balance.hashed_key}"
},
json={
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
},
)
assert response.status_code == 500
error_data = response.json()
assert error_data["error"]["type"] == "internal_error"
assert error_data["error"]["message"] == "An unexpected server error occurred"
@pytest.mark.asyncio
async def test_proxy_with_model_based_pricing(
async_client: AsyncClient, test_session: AsyncSession
) -> None:
"""Test proxy with model-based pricing enabled."""
# Create API key with sufficient balance
unique_id = str(uuid.uuid4())[:8]
key = ApiKey(
hashed_key=f"model-pricing-key-{unique_id}",
balance=10000000, # 10,000 sats
refund_address=None,
total_spent=0,
total_requests=0,
)
test_session.add(key)
await test_session.commit()
with patch.dict(os.environ, {"MODEL_BASED_PRICING": "true"}):
with patch("os.path.exists", return_value=True):
# Mock a model with pricing
from router.models import MODELS, Architecture, Model, Pricing, TopProvider
test_model = Model(
id="gpt-4",
name="GPT-4",
created=1680000000,
description="Test model",
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="cl100k_base",
instruct_type="none",
),
pricing=Pricing(
prompt=0.03,
completion=0.06,
request=0.001,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
),
sats_pricing=Pricing(
prompt=300, # 300 sats per 1k tokens
completion=600,
request=10,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
max_cost=5000, # 5000 sats max
),
top_provider=TopProvider(
context_length=8192, max_completion_tokens=4096, is_moderated=False
),
)
# Temporarily replace models
original_models = MODELS[:]
MODELS.clear()
MODELS.append(test_model)
# Mock the upstream HTTP client
with patch("httpx.AsyncClient") as mock_client_class:
mock_client = AsyncMock()
mock_client_class.return_value = mock_client
# Create a mock response
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.aread = AsyncMock(
return_value=b'{"id": "test", "model": "gpt-4"}'
)
mock_response.aiter_bytes = AsyncMock()
mock_response.aclose = AsyncMock()
mock_client.send = AsyncMock(return_value=mock_response)
mock_client.build_request = AsyncMock()
mock_client.aclose = AsyncMock()
try:
response = await async_client.post(
"/v1/chat/completions",
headers={"Authorization": f"Bearer sk-{key.hashed_key}"},
json={
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
},
)
# Should succeed because balance (10,000 sats) > max_cost (5000 sats)
assert response.status_code == 200
finally:
MODELS.clear()
MODELS.extend(original_models)
Generated
+1703 -1263
View File
File diff suppressed because it is too large Load Diff