mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge branch 'main' into kwsantiago/62-comprehensive-tests
This commit is contained in:
+17
-25
@@ -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
|
||||
|
||||
@@ -1,5 +0,0 @@
|
||||
- test if currency and payment amount is correct
|
||||
- test payout
|
||||
|
||||
- make tor work
|
||||
-
|
||||
+2
-1
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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__)
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -0,0 +1,3 @@
|
||||
from .logging import get_logger
|
||||
|
||||
__all__ = ["get_logger"]
|
||||
@@ -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()">×</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}
|
||||
@@ -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)
|
||||
@@ -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
@@ -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}
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost
|
||||
|
||||
__all__ = [
|
||||
"CostData",
|
||||
"CostDataError",
|
||||
"MaxCostData",
|
||||
"calculate_cost",
|
||||
]
|
||||
|
||||
@@ -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
@@ -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}
|
||||
@@ -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
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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
@@ -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,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")
|
||||
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user