Merge pull request #83 from Routstr/cleanup

Cleanup and smaller fixes
This commit is contained in:
shroominic
2025-08-02 14:58:32 -03:00
committed by GitHub
26 changed files with 700 additions and 281 deletions
+11 -25
View File
@@ -1,5 +1,4 @@
NAME = "Your Routstr Proxy Name"
DESCRIPTION = "A short Description"
# Any openai-compatible api endpoint
@@ -9,40 +8,27 @@ UPSTREAM_API_KEY="sk-21212121212121212121212121212121"
# 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"
# 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"
-5
View File
@@ -1,5 +0,0 @@
- test if currency and payment amount is correct
- test payout
- make tor work
-
+104 -14
View File
@@ -1,11 +1,12 @@
{
"models": [
{
"id": "google/gemini-2.5-pro-preview",
"id": "google/gemini-2.5-flash",
"canonical_slug": "google/gemini-2.5-flash",
"hugging_face_id": "",
"name": "Google: Gemini 2.5 Pro Preview 06-05",
"created": 1749137257,
"description": "Gemini 2.5 Pro is Google\u2019s state-of-the-art AI model designed for advanced reasoning, coding, mathematics, and scientific tasks. It employs \u201cthinking\u201d capabilities, enabling it to reason through responses with enhanced accuracy and nuanced context handling. Gemini 2.5 Pro achieves top-tier performance on multiple benchmarks, including first-place positioning on the LMArena leaderboard, reflecting superior human-preference alignment and complex problem-solving abilities.\n",
"name": "Google: Gemini 2.5 Flash",
"created": 1750172488,
"description": "Gemini 2.5 Flash is Google's state-of-the-art workhorse model, specifically designed for advanced reasoning, coding, mathematics, and scientific tasks. It includes built-in \"thinking\" capabilities, enabling it to provide responses with greater accuracy and nuanced context handling. \n\nAdditionally, Gemini 2.5 Flash is configurable through the \"max tokens for reasoning\" parameter, as described in the documentation (https://openrouter.ai/docs/use-cases/reasoning-tokens#max-tokens-for-reasoning).",
"context_length": 1048576,
"architecture": {
"modality": "text+image->text",
@@ -21,18 +22,107 @@
"instruct_type": null
},
"pricing": {
"prompt": "0.00000125",
"completion": "0.00001",
"prompt": "0.0000003",
"completion": "0.0000025",
"request": "0",
"image": "0.00516",
"image": "0.001238",
"web_search": "0",
"internal_reasoning": "0",
"input_cache_read": "0.00000031",
"input_cache_write": "0.000001625"
"input_cache_read": "0.000000075",
"input_cache_write": "0.0000003833"
},
"top_provider": {
"context_length": 1048576,
"max_completion_tokens": 65536,
"max_completion_tokens": 65535,
"is_moderated": false
},
"per_request_limits": null,
"supported_parameters": [
"max_tokens",
"temperature",
"top_p",
"tools",
"tool_choice",
"stop",
"response_format",
"structured_outputs"
]
},
{
"id": "openai/o3-pro",
"canonical_slug": "openai/o3-pro-2025-06-10",
"hugging_face_id": "",
"name": "OpenAI: o3 Pro",
"created": 1749598352,
"description": "The o-series of models are trained with reinforcement learning to think before they answer and perform complex reasoning. The o3-pro model uses more compute to think harder and provide consistently better answers.\n\nNote that BYOK is required for this model. Set up here: https://openrouter.ai/settings/integrations",
"context_length": 200000,
"architecture": {
"modality": "text+image->text",
"input_modalities": [
"text",
"file",
"image"
],
"output_modalities": [
"text"
],
"tokenizer": "Other",
"instruct_type": null
},
"pricing": {
"prompt": "0.00002",
"completion": "0.00008",
"request": "0",
"image": "0.0153",
"web_search": "0",
"internal_reasoning": "0"
},
"top_provider": {
"context_length": 200000,
"max_completion_tokens": 100000,
"is_moderated": true
},
"per_request_limits": null,
"supported_parameters": [
"tools",
"tool_choice",
"seed",
"max_tokens",
"response_format",
"structured_outputs"
]
},
{
"id": "x-ai/grok-3-mini",
"canonical_slug": "x-ai/grok-3-mini",
"hugging_face_id": "",
"name": "xAI: Grok 3 Mini",
"created": 1749583245,
"description": "A lightweight model that thinks before responding. Fast, smart, and great for logic-based tasks that do not require deep domain knowledge. The raw thinking traces are accessible.",
"context_length": 131072,
"architecture": {
"modality": "text->text",
"input_modalities": [
"text"
],
"output_modalities": [
"text"
],
"tokenizer": "Grok",
"instruct_type": null
},
"pricing": {
"prompt": "0.0000003",
"completion": "0.0000005",
"request": "0",
"image": "0",
"web_search": "0",
"internal_reasoning": "0",
"input_cache_read": "0.000000075"
},
"top_provider": {
"context_length": 131072,
"max_completion_tokens": null,
"is_moderated": false
},
"per_request_limits": null,
@@ -45,11 +135,11 @@
"reasoning",
"include_reasoning",
"structured_outputs",
"response_format",
"stop",
"frequency_penalty",
"presence_penalty",
"seed"
"seed",
"logprobs",
"top_logprobs",
"response_format"
]
}
]
+1 -1
View File
@@ -1,6 +1,6 @@
[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"
+1 -1
View File
@@ -2,6 +2,6 @@ import dotenv
dotenv.load_dotenv()
from .main import app as fastapi_app # noqa
from .core.main import app as fastapi_app # noqa
__all__ = ["fastapi_app"]
+2 -2
View File
@@ -4,8 +4,8 @@ from typing import Optional
from fastapi import HTTPException
from sqlmodel import col, update
from .db import ApiKey, AsyncSession
from .logging import get_logger
from .core import get_logger
from .core.db import ApiKey, AsyncSession
from .payment.cost_caculation import (
CostData,
CostDataError,
+13 -7
View File
@@ -3,10 +3,11 @@ from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException
from .auth import validate_bearer_key
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(
@@ -23,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,
@@ -31,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,
@@ -39,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),
@@ -49,7 +50,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),
@@ -82,7 +83,7 @@ async def refund_wallet_endpoint(
return result
@wallet_router.api_route(
@router.api_route(
"/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"],
include_in_schema=False,
@@ -92,3 +93,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)
+3
View File
@@ -0,0 +1,3 @@
from .logging import get_logger
__all__ = ["get_logger"]
+2 -2
View File
@@ -6,10 +6,10 @@ 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
from .wallet import get_balance, send_token
admin_router = APIRouter(prefix="/admin")
admin_router = APIRouter(prefix="/admin", include_in_schema=False)
class WithdrawRequest(BaseModel):
View File
+13 -8
View File
@@ -8,6 +8,7 @@ from pathlib import Path
from typing import Any
from pythonjsonlogger import jsonlogger
from rich.logging import RichHandler
class DailyRotatingFileHandler(logging.handlers.TimedRotatingFileHandler):
@@ -181,10 +182,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},
@@ -192,11 +189,13 @@ def setup_logging() -> None:
},
"handlers": {
"console": {
"class": "logging.StreamHandler",
"()": RichHandler,
"level": log_level,
"formatter": "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,
@@ -252,6 +251,12 @@ def setup_logging() -> None:
"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,
+10 -26
View File
@@ -6,14 +6,14 @@ from typing import AsyncGenerator
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from .account import wallet_router
from ..balance import balance_router, deprecated_wallet_router
from ..discovery import providers_router
from ..payment.models import MODELS, models_router, update_sats_pricing
from ..proxy import proxy_router
from ..wallet import check_for_refunds, periodic_payout
from .admin import admin_router
from .db import init_db
from .discovery import providers_router
from .logging import get_logger, setup_logging
from .models import MODELS, models_router, update_sats_pricing
from .proxy import proxy_router
from .wallet import check_for_refunds, periodic_payout
# Initialize logging first
setup_logging()
@@ -28,19 +28,11 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
try:
await init_db()
logger.info("Database initialized successfully")
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:
@@ -85,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 {
@@ -99,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,
@@ -108,11 +96,7 @@ async def info() -> dict:
app.include_router(models_router)
app.include_router(admin_router)
app.include_router(wallet_router)
app.include_router(balance_router)
app.include_router(deprecated_wallet_router)
app.include_router(providers_router)
app.include_router(proxy_router)
logger.info(
"Application initialized successfully",
extra={"version": __version__, "routers_count": 5},
)
+163 -102
View File
@@ -2,8 +2,8 @@ import asyncio
import json
import os
import random
import re
import string
from typing import Any
import httpx
import websockets
@@ -17,59 +17,34 @@ def generate_subscription_id() -> str:
return "".join(random.choices(string.ascii_lowercase + string.digits, k=10))
def extract_onion_urls(content: str) -> list[str]:
"""Extract onion URLs from content."""
pattern = r"http?://[a-zA-Z0-9\-._~]+\.onion"
return re.findall(pattern, content)
async def query_nostr_relay_with_search(
search_term: str,
async def query_nostr_relay_for_providers(
relay_url: str,
kinds: list[int] | None = None,
pubkey: str | None = None,
limit: int = 1000,
timeout: int = 30,
) -> list[dict]:
) -> list[dict[str, Any]]:
"""
Query a Nostr relay and filter for events containing a search term.
Query a Nostr relay for provider announcements using RIP-02 spec.
Searches for kind 31338 events (Routstr Provider Announcements).
"""
if kinds is None:
kinds = [1]
events = []
# If searching for an npub mention, try tag-based search first
if search_term.startswith("nostr:npub"):
# Extract the npub and convert to hex
npub = search_term.replace("nostr:", "")
try:
# Convert npub to hex (you might need to implement or import this)
# For now, try tag-based search with the npub
filter_obj = {
"kinds": kinds,
"limit": limit,
"#p": [npub], # Posts that tag this pubkey
}
except Exception:
# If conversion fails, try regular search
filter_obj = {
"kinds": kinds,
"limit": limit,
}
else:
# Try relay's search functionality (NIP-50)
filter_obj = {
"kinds": kinds,
"search": search_term,
"limit": limit,
}
# Build filter according to RIP-02 spec
filter_obj: dict[str, Any] = {
"kinds": [31338], # RIP-02 Provider Announcement events
"limit": limit,
}
# If specific pubkey provided, filter by author
if pubkey:
filter_obj["authors"] = [pubkey]
sub_id = generate_subscription_id()
req_message = json.dumps(["REQ", sub_id, filter_obj])
try:
async with websockets.connect(relay_url, timeout=timeout) as websocket:
print(f"Connected to relay, sending request with filter: {filter_obj}")
print("Connected to relay, searching for kind 31338 events")
await websocket.send(req_message)
while True:
@@ -78,25 +53,14 @@ async def query_nostr_relay_with_search(
data = json.loads(message)
if data[0] == "EVENT" and data[1] == sub_id:
# For tag-based search, also check content
if search_term.startswith("nostr:npub"):
if search_term.lower() in data[2]["content"].lower():
print(f"Found matching event: {data[2]['id']}")
events.append(data[2])
else:
print(f"Found matching event: {data[2]['id']}")
events.append(data[2])
event = data[2]
print(f"Found provider announcement: {event['id']}")
events.append(event)
elif data[0] == "EOSE" and data[1] == sub_id:
print("Received EOSE message")
break
elif data[0] == "NOTICE":
print(f"Relay notice: {data[1]}")
# If search not supported, could break and try different approach
if "unrecognised filter item" in data[1] and "search" in str(
filter_obj
):
print("Search not supported on this relay")
break
except asyncio.TimeoutError:
print("Timeout waiting for message")
@@ -110,61 +74,159 @@ async def query_nostr_relay_with_search(
except Exception as e:
print(f"Query failed: {e}")
print(f"Query complete. Found {len(events)} matching events")
print(f"Query complete. Found {len(events)} provider announcements")
return events
async def get_cache() -> list[dict]:
def parse_provider_announcement(event: dict[str, Any]) -> dict[str, Any] | None:
"""
Parse a kind 31338 provider announcement event according to RIP-02 spec.
Returns structured provider data or None if invalid.
"""
try:
# Extract required tags according to RIP-02
tags = event.get("tags", [])
# Find required tags
endpoint_url = None
provider_name = None
d_tag = None
for tag in tags:
if len(tag) >= 2:
if tag[0] == "endpoint":
endpoint_url = tag[1]
elif tag[0] == "name":
provider_name = tag[1]
elif tag[0] == "d":
d_tag = tag[1]
# Validate required fields
if not endpoint_url or not provider_name or not d_tag:
print(
f"Invalid provider announcement - missing required tags: {event['id']}"
)
return None
# Extract optional tags
description = None
contact = None
pricing_url = None
supported_models = []
for tag in tags:
if len(tag) >= 2:
if tag[0] == "description":
description = tag[1]
elif tag[0] == "contact":
contact = tag[1]
elif tag[0] == "pricing":
pricing_url = tag[1]
elif tag[0] == "model":
supported_models.append(tag[1])
return {
"id": event["id"],
"pubkey": event["pubkey"],
"created_at": event["created_at"],
"d_tag": d_tag,
"endpoint_url": endpoint_url,
"name": provider_name,
"description": description,
"contact": contact,
"pricing_url": pricing_url,
"supported_models": supported_models,
"content": event.get("content", ""),
}
except Exception as e:
print(f"Error parsing provider announcement {event.get('id', 'unknown')}: {e}")
return None
async def get_cache() -> list[dict[str, Any]]:
return [] # TODO: Implement cache
async def fetch_onion(provider: str) -> dict:
"""Check if an onion service is healthy by making a GET request to its root."""
async def fetch_provider_health(endpoint_url: str) -> dict[str, Any]:
"""Check if a provider endpoint is healthy by making a GET request."""
try:
# Get Tor proxy URL from environment variable, default to local Tor SOCKS5 proxy
tor_proxy = os.getenv("TOR_PROXY_URL", "socks5://127.0.0.1:9050")
# Determine if we need Tor proxy based on .onion domain
is_onion = ".onion" in endpoint_url
# Set up client arguments conditionally
proxies = None
if is_onion:
# Get Tor proxy URL from environment variable
tor_proxy = os.getenv("TOR_PROXY_URL", "socks5://127.0.0.1:9050")
proxies = {"http://": tor_proxy, "https://": tor_proxy} # type: ignore[assignment]
# Configure httpx to use Tor SOCKS5 proxy
async with httpx.AsyncClient(
proxies={"http://": tor_proxy, "https://": tor_proxy}, # type: ignore
timeout=httpx.Timeout(30.0),
follow_redirects=True,
proxies=proxies, # type: ignore[arg-type]
) as client:
response = await client.get(provider)
# Consider 2xx and 3xx status codes as healthy
return {"status_code": response.status_code, "json": response.json()}
except Exception:
# Any exception means the service is not healthy
return {"status_code": 500, "json": {"error": "Failed to fetch onion"}}
# Try to fetch models endpoint first (common for AI providers)
models_url = f"{endpoint_url.rstrip('/')}/v1/models"
try:
response = await client.get(models_url)
if response.status_code == 200:
return {
"status_code": response.status_code,
"endpoint": "models",
"json": response.json(),
}
except Exception:
pass
# Fallback to root endpoint
response = await client.get(endpoint_url)
return {
"status_code": response.status_code,
"endpoint": "root",
"json": response.json()
if response.headers.get("content-type", "").startswith(
"application/json"
)
else {"message": "OK"},
}
except Exception as e:
return {
"status_code": 500,
"endpoint": "error",
"json": {"error": f"Failed to fetch provider: {str(e)}"},
}
@providers_router.get("/")
async def get_providers(include_json: bool = False) -> dict[str, list[dict | str]]:
npub = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s"
async def get_providers(
include_json: bool = False, pubkey: str | None = None
) -> dict[str, list[dict[str, Any]]]:
"""
Discover Routstr providers using RIP-02 specification.
Searches for kind 31338 provider announcement events on Nostr relays.
# Relays that support NIP-50 text search
search_relays = [
"wss://relay.nostr.band", # Known to support search
"wss://nostr.wine", # Known to support search
Reference: https://github.com/Routstr/protocol/blob/main/RIP-02.md
"""
# Default relays for provider discovery
discovery_relays = [
"wss://relay.nostr.band",
"wss://relay.damus.io",
"wss://nos.lol",
"wss://relay.routstr.com",
]
# Search for the mention format that appears in posts
search_term = f"nostr:{npub}"
all_events = []
event_ids = set() # To avoid duplicates
# Try multiple relays
for relay_url in search_relays:
print(f"\nTrying relay: {relay_url}")
# Query multiple relays for provider announcements
for relay_url in discovery_relays:
print(f"\nQuerying relay for providers: {relay_url}")
try:
events = await query_nostr_relay_with_search(
search_term=search_term,
events = await query_nostr_relay_for_providers(
relay_url=relay_url,
kinds=[1], # Text notes
limit=500,
pubkey=pubkey,
limit=100,
)
# Add unique events
@@ -173,35 +235,34 @@ async def get_providers(include_json: bool = False) -> dict[str, list[dict | str
event_ids.add(event["id"])
all_events.append(event)
print(f"Got {len(events)} events from {relay_url}")
# If we have enough events, we can stop
if len(all_events) >= 100:
break
print(f"Got {len(events)} provider announcements from {relay_url}")
except Exception as e:
print(f"Failed to query {relay_url}: {e}")
continue
print(f"Found {len(all_events)} total unique events mentioning routstr")
print(f"Found {len(all_events)} total unique provider announcements")
# Parse provider announcements according to RIP-02
providers = []
for event in all_events:
onion_urls = extract_onion_urls(event["content"])
providers.extend(onion_urls)
parsed_provider = parse_provider_announcement(event)
if parsed_provider:
providers.append(parsed_provider)
unique_providers = list(set(providers))
print(f"Parsed {len(providers)} valid provider announcements")
print(f"Found {len(unique_providers)} unique onion URLs")
print(unique_providers)
healthy_providers: list[dict | str] = []
for provider in unique_providers:
response = await fetch_onion(provider)
# Check provider health if requested
healthy_providers: list[dict[str, Any]] = []
for provider in providers:
endpoint_url = provider["endpoint_url"]
if include_json:
healthy_providers.append({provider: response["json"]})
health_check = await fetch_provider_health(endpoint_url)
provider_data = {"provider": provider, "health": health_check}
healthy_providers.append(provider_data)
else:
# Just return the provider info without health check
healthy_providers.append(provider)
return {"providers": healthy_providers}
+8
View File
@@ -0,0 +1,8 @@
from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost
__all__ = [
"CostData",
"CostDataError",
"MaxCostData",
"calculate_cost",
]
+2 -2
View File
@@ -2,8 +2,8 @@ import os
from pydantic import BaseModel
from ..logging import get_logger
from ..models import MODELS
from ..core import get_logger
from .models import MODELS
logger = get_logger(__name__)
+4 -4
View File
@@ -6,9 +6,9 @@ from typing import Literal
import cbor2
from fastapi import HTTPException, Response
from ..logging import get_logger
from ..models import MODELS
from ..core import get_logger
from .cost_caculation import COST_PER_REQUEST, MODEL_BASED_PRICING
from .models import MODELS
logger = get_logger(__name__)
@@ -113,7 +113,7 @@ def check_token_balance(headers: dict, body: dict) -> Literal["sat", "msat"]:
try:
_token = base64_token_json(cashu_token)
amount = sum(p["amount"] for t in _token["token"] for p in t["proofs"])
unit: Literal["sat", "msat"] = _token["unit"]
unit: Literal["sat", "msat"] = _token.get("unit", "sat")
if unit == "sat":
amount *= 1000
@@ -122,7 +122,7 @@ def check_token_balance(headers: dict, body: dict) -> Literal["sat", "msat"]:
"CashuA token parsed successfully",
extra={
"amount": amount,
"unit": _token["unit"],
"unit": unit,
"amount_msats": amount,
"required_cost_msats": cost,
"sufficient_balance": amount >= cost,
@@ -144,7 +144,6 @@ async def update_sats_pricing() -> None:
break
@models_router.get("/models")
@models_router.get("/v1/models")
async def models() -> dict:
return {"data": MODELS}
+1 -1
View File
@@ -3,7 +3,7 @@ import os
import httpx
from .logging import get_logger
from ..core import get_logger
logger = get_logger(__name__)
+43 -6
View File
@@ -6,9 +6,14 @@ import httpx
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import Response, StreamingResponse
from ..logging import get_logger
from ..core import get_logger
from ..wallet import CurrencyUnit, recieve_token, send_token
from .cost_caculation import CostData, CostDataError, MaxCostData, calculate_cost
from .cost_caculation import (
CostData,
CostDataError,
MaxCostData,
calculate_cost,
)
from .helpers import (
UPSTREAM_BASE_URL,
create_error_response,
@@ -46,16 +51,37 @@ async def x_cashu_handler(
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(
@@ -297,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",
@@ -399,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",
+2 -2
View File
@@ -13,8 +13,8 @@ from .auth import (
revert_pay_for_request,
validate_bearer_key,
)
from .db import ApiKey, AsyncSession, create_session, get_session
from .logging import get_logger
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,
+38 -60
View File
@@ -1,24 +1,11 @@
import os
from typing import Literal
from cashu.core.base import Token
from cashu.core.base import Token, Unit
from cashu.wallet.helpers import deserialize_token_from_string, receive, send
from cashu.wallet.wallet import Wallet
from .db import ApiKey, AsyncSession
from .logging import get_logger
# from .cashu import (
# credit_balance,
# delete_key_if_zero_balance,
# refund_balance,
# wallet,
# )
# from .cashu import (
# check_for_refunds,
# init_wallet,
# periodic_payout,
# )
from .core import db, get_logger
logger = get_logger(__name__)
@@ -51,10 +38,12 @@ async def recieve_token(
load_all_keysets=True,
unit=token_obj.unit,
)
if token_obj.mint in TRUSTED_MINTS and token_obj.mint != PRIMARY_MINT_URL:
return await swap_to_primary_mint(token_obj, wallet)
elif token_obj.mint not in TRUSTED_MINTS:
raise ValueError("Mint URL is not supported by this proxy")
if token_obj.mint != PRIMARY_MINT_URL:
raise ValueError(
f"This mint is not supported, please use {PRIMARY_MINT_URL} instead"
)
await receive(wallet, token_obj)
return token_obj.amount, token_obj.unit, token_obj.mint
@@ -73,9 +62,16 @@ async def send_token(
async def swap_to_primary_mint(
token_obj: Token, wallet: Wallet
token_obj: Token, token_wallet: Wallet
) -> tuple[int, CurrencyUnit, str]:
print(f"swap_to_primary_mint, token_obj: {token_obj}")
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":
@@ -84,46 +80,33 @@ async def swap_to_primary_mint(
raise ValueError("Invalid unit")
estimated_fee_sat = max(amount_msat // 1000 * 0.01, 2)
amount_msat_after_fee = amount_msat - estimated_fee_sat * 1000
mint_quote = await wallet.mint_quote(amount_msat_after_fee, "sat")
melt_quote = await wallet.melt_quote(mint_quote.request, amount_msat_after_fee)
_ = await wallet.melt(
print(f"amount_msat_after_fee: {amount_msat_after_fee}")
primary_wallet = await Wallet.with_db(
PRIMARY_MINT_URL, db=".temp", load_all_keysets=True, unit="sat"
)
await primary_wallet.load_mint_keysets()
mint_quote = await primary_wallet.mint_quote(
amount_msat_after_fee // 1000, Unit.sat
)
print(f"mint_quote: {mint_quote}")
melt_quote = await token_wallet.melt_quote(mint_quote.request)
print(f"melt_quote: {melt_quote}")
melt_quote_resp = await token_wallet.melt(
proofs=token_obj.proofs,
invoice=mint_quote.request,
fee_reserve=melt_quote.fee_reserve,
fee_reserve_sat=melt_quote.fee_reserve,
quote_id=melt_quote.quote,
)
print(f"melt_quote_resp: {melt_quote_resp}")
_ = await wallet.mint(token_obj.amount, mint_quote.quote)
_ = await primary_wallet.mint(amount_msat_after_fee // 1000, mint_quote.quote)
return token_obj.amount, "sat", PRIMARY_MINT_URL
return amount_msat_after_fee // 1000, "sat", PRIMARY_MINT_URL
# insert initial token state here to reduce db calls
# async def create_refund_token(
# amount: int, unit: CurrencyUnit, mint_url: str | None = None
# ) -> str:
# wallet = await Wallet.with_db(
# mint_url, DATABASE_URL, load_all_keysets=True, unit=unit
# )
# if wallet.balance_per_minturl(unit=unit)[mint_url] < amount:
# raise ValueError("Wallet has no balance")
# if mint_url is None:
# mint_url = wallet.mint_urls[0]
# return await wallet._make_token(amount, unit=unit, mint_url=mint_url)
# async def redeem_token(token: str) -> Token:
# token_obj = deserialize_token_from_string(token)
# wallet = await Wallet.with_db(
# token_obj.mint,
# DATABASE_URL,
# load_all_keysets=True,
# unit=token_obj.unit,
# )
# return await redeem_universal(wallet, token_obj)
async def credit_balance(cashu_token: str, key: ApiKey, session: AsyncSession) -> int:
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
@@ -144,16 +127,11 @@ async def send_to_lnurl(amount: int, unit: CurrencyUnit, lnurl: str) -> dict[str
async def check_for_refunds() -> None:
print("check_for_refunds, temp not implemented")
async def init_wallet() -> None:
balance = await get_balance("sat")
print(f"init_wallet, balance: {balance}")
logger.warning("check_for_refunds, temporary not implemented")
async def periodic_payout() -> None:
print("periodic_payout, temp not implemented")
logger.warning("periodic_payout, temporary not implemented")
# class Proof:
+268
View File
@@ -0,0 +1,268 @@
#!/usr/bin/env python3
"""
Simple Python function to publish one provider listing to a nostr relay
according to the RIP-02 specification.
Based on: https://github.com/Routstr/protocol/blob/main/RIP-02.md
Event Kind: 31338 (Routstr Provider Announcements)
"""
import asyncio
import hashlib
import json
import time
from typing import Any
import secp256k1
import websockets
def create_provider_announcement_event(
private_key_hex: str,
provider_name: str,
endpoint_url: str,
d_tag: str,
description: str | None = None,
contact: str | None = None,
pricing_url: str | None = None,
supported_models: list[str] | None = None,
) -> dict[str, Any]:
"""
Create a RIP-02 compliant provider announcement event.
Args:
private_key_hex: 32-byte hex private key for signing
provider_name: Human readable name for the provider
endpoint_url: Base URL for the provider's API endpoint
d_tag: Unique identifier for this provider (required for addressable events)
description: Optional description of the provider
contact: Optional contact information
pricing_url: Optional URL to pricing information
supported_models: Optional list of supported model names
Returns:
Complete signed nostr event ready for publishing
"""
# Convert hex private key to secp256k1 PrivateKey object
private_key = secp256k1.PrivateKey(bytes.fromhex(private_key_hex))
public_key = private_key.pubkey.serialize(compressed=True)[
1:
] # Remove 0x02/0x03 prefix
# Build required tags according to RIP-02
tags = [
["d", d_tag], # Required for addressable events (kind 30000-39999)
["endpoint", endpoint_url],
["name", provider_name],
]
# Add optional tags if provided
if description:
tags.append(["description", description])
if contact:
tags.append(["contact", contact])
if pricing_url:
tags.append(["pricing", pricing_url])
if supported_models:
for model in supported_models:
tags.append(["model", model])
# Create the event structure
created_at = int(time.time())
event_data = [
0, # Reserved field
public_key.hex(), # Public key as hex
created_at, # Unix timestamp
31338, # Kind for RIP-02 Provider Announcements
tags, # Tags array
"", # Content (empty for provider announcements)
]
# Serialize event data for hashing
event_json = json.dumps(event_data, separators=(",", ":"), ensure_ascii=False)
# Calculate event ID (SHA256 hash)
event_id = hashlib.sha256(event_json.encode("utf-8")).hexdigest()
# Sign the event ID
signature = private_key.ecdsa_sign(bytes.fromhex(event_id), raw=True)
signature_der = private_key.ecdsa_serialize(signature)
# Create the final event
event = {
"id": event_id,
"pubkey": public_key.hex(),
"created_at": created_at,
"kind": 31338,
"tags": tags,
"content": "",
"sig": signature_der.hex(),
}
return event
async def publish_provider_to_relay(
relay_url: str, event: dict[str, Any], timeout: int = 30
) -> bool:
"""
Publish a provider announcement event to a nostr relay.
Args:
relay_url: WebSocket URL of the nostr relay (e.g., "wss://relay.damus.io")
event: Complete signed nostr event to publish
timeout: Connection timeout in seconds
Returns:
True if successfully published, False otherwise
"""
try:
async with websockets.connect(relay_url, timeout=timeout) as websocket:
# Send EVENT message
event_message = json.dumps(["EVENT", event])
await websocket.send(event_message)
print(f"Published event {event['id']} to {relay_url}")
# Wait for OK response
try:
response = await asyncio.wait_for(websocket.recv(), timeout=5)
data = json.loads(response)
if data[0] == "OK" and data[1] == event["id"]:
if data[2]: # True means accepted
print(
f"✅ Event accepted by relay: {data[3] if len(data) > 3 else ''}"
)
return True
else:
print(
f"❌ Event rejected by relay: {data[3] if len(data) > 3 else ''}"
)
return False
elif data[0] == "NOTICE":
print(f"📢 Relay notice: {data[1]}")
return False
else:
print(f"🤔 Unexpected response: {data}")
return False
except asyncio.TimeoutError:
print("⏰ No response from relay within timeout")
return False
except Exception as e:
print(f"💥 Failed to publish to {relay_url}: {e}")
return False
async def publish_provider_listing(
private_key_hex: str,
provider_name: str,
endpoint_url: str,
d_tag: str,
relay_urls: list[str] | None = None,
description: str | None = None,
contact: str | None = None,
pricing_url: str | None = None,
supported_models: list[str] | None = None,
) -> dict[str, bool]:
"""
Complete function to create and publish a provider listing to nostr relays.
Args:
private_key_hex: 32-byte hex private key for signing
provider_name: Human readable name for the provider
endpoint_url: Base URL for the provider's API endpoint
d_tag: Unique identifier for this provider
relay_urls: List of relay URLs to publish to (uses defaults if None)
description: Optional description of the provider
contact: Optional contact information
pricing_url: Optional URL to pricing information
supported_models: Optional list of supported model names
Returns:
Dictionary mapping relay URLs to success status
"""
# Use default relays if none provided
if relay_urls is None:
relay_urls = [
"wss://relay.nostr.band",
"wss://relay.damus.io",
"wss://relay.routstr.com",
]
# Create the provider announcement event
event = create_provider_announcement_event(
private_key_hex=private_key_hex,
provider_name=provider_name,
endpoint_url=endpoint_url,
d_tag=d_tag,
description=description,
contact=contact,
pricing_url=pricing_url,
supported_models=supported_models,
)
print(f"📝 Created provider announcement event: {event['id']}")
print(f"🔑 Public key: {event['pubkey']}")
print(f"🏷️ Provider: {provider_name}")
print(f"🌐 Endpoint: {endpoint_url}")
print()
# Publish to all specified relays
results = {}
tasks = []
for relay_url in relay_urls:
task = publish_provider_to_relay(relay_url, event)
tasks.append((relay_url, task))
# Execute all publishing tasks concurrently
for relay_url, task in tasks:
try:
success = await task
results[relay_url] = success
except Exception as e:
print(f"💥 Failed to publish to {relay_url}: {e}")
results[relay_url] = False
return results
# Example usage
async def main() -> None:
"""Example of how to use the provider publishing function."""
# Example private key (DO NOT use this in production!)
private_key = "3185a47e3802f956ca207b46c8d6b8b5c5dbad53a5ca29816050e9b66badc33c"
# Example provider information
provider_name = "My AI Provider"
endpoint_url = "https://api.myaiprovider.com"
d_tag = "my-ai-provider-v1" # Unique identifier
description = "High-quality AI models with competitive pricing"
contact = "admin@myaiprovider.com"
pricing_url = "https://myaiprovider.com/pricing"
supported_models = ["gpt-4o", "claude-3-sonnet", "llama-3.1-70b"]
# Publish to relays
results = await publish_provider_listing(
private_key_hex=private_key,
provider_name=provider_name,
endpoint_url=endpoint_url,
d_tag=d_tag,
description=description,
contact=contact,
pricing_url=pricing_url,
supported_models=supported_models,
)
# Print results
print("\n📊 Publishing Results:")
for relay_url, success in results.items():
status = "✅ Success" if success else "❌ Failed"
print(f" {relay_url}: {status}")
if __name__ == "__main__":
asyncio.run(main())
+1 -2
View File
@@ -28,9 +28,8 @@ To run specific test files:
```bash
pytest tests/test_main.py
pytest tests/test_account.py
pytest tests/test_proxy.py
pytest tests/test_models.py
pytest tests/test_proxy.py
```
To run only async tests:
+4 -4
View File
@@ -37,8 +37,8 @@ TEST_ENV = {
os.environ.update(TEST_ENV)
# Now import modules that depend on environment variables
from router.db import get_session # noqa: E402
from router.main import app # noqa: E402
from router.core.db import get_session # noqa: E402
from router.core.main import app # noqa: E402
@pytest.fixture(scope="session")
@@ -79,7 +79,7 @@ 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("router.models.update_sats_pricing") as mock_update:
with patch("router.payment.models.update_sats_pricing") as mock_update:
mock_update.return_value = None
yield TestClient(app)
@@ -95,7 +95,7 @@ async def async_client(test_session: AsyncSession) -> AsyncGenerator[AsyncClient
# Mock startup tasks
with patch.dict(os.environ, TEST_ENV, clear=True):
with patch("router.models.update_sats_pricing") as mock_update:
with patch("router.payment.models.update_sats_pricing") as mock_update:
mock_update.return_value = None
async with AsyncClient(
+2 -2
View File
@@ -13,7 +13,7 @@ async def test_root_endpoint(async_client: AsyncClient) -> None:
"NAME": "TestRoutstrNode",
"DESCRIPTION": "Test Node",
"NPUB": "npub1test",
"MINT": "https://test.mint.com",
"CASHU_MINTS": "https://test.mint.com,https://test.mint2.com",
"HTTP_URL": "http://test.example.com",
"ONION_URL": "http://test.onion",
}
@@ -28,7 +28,7 @@ async def test_root_endpoint(async_client: AsyncClient) -> None:
assert "name" in data
assert "description" in data
assert "npub" in data
assert "mint" in data
assert "mints" in data
assert "http_url" in data
assert "onion_url" in data
+4 -4
View File
@@ -4,7 +4,7 @@ from unittest.mock import AsyncMock, patch
import pytest
from router.models import (
from router.payment.models import (
MODELS,
Architecture,
Model,
@@ -49,7 +49,7 @@ async def test_update_sats_pricing_calculation(sample_model: Model) -> None:
"""Test that sats pricing is calculated correctly."""
# Mock the sats_usd_ask_price function
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
"router.payment.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD
@@ -137,7 +137,7 @@ async def test_update_sats_pricing_without_top_provider() -> None:
)
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
"router.payment.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD
@@ -189,7 +189,7 @@ async def test_update_sats_pricing_without_top_provider() -> None:
async def test_update_sats_pricing_handles_errors() -> None:
"""Test that update_sats_pricing handles errors gracefully."""
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
"router.payment.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.side_effect = Exception("API Error")