Merge branch 'main' into bug-fixes

This commit is contained in:
shroominic
2025-05-28 13:10:02 +02:00
committed by GitHub
10 changed files with 5347 additions and 951 deletions
+2 -2
View File
@@ -3,9 +3,9 @@ __pycache__
keys.db
wallet.sqlite3
# Development
.notes
.keys.db
.wallet.sqlite3
.models.json
.models.json
compose.override.yml
+1 -9
View File
@@ -1,12 +1,4 @@
# proxy
a reverse proxy that you can plug in front of any openai compatible api endpoint
to handle api-key based payments using cashu tokens or bold12 lightning invoices
we also want to provice an internal dashboard
- general settings
- configure payment methods
- request evals
- publish your listing to nostr
- monitor traffic
to handle payments using the cashu protocol (Bitcoin L3)
+3 -2
View File
@@ -7,9 +7,10 @@ services:
- .:/app
env_file:
- .env
environment:
- TOR_PROXY_URL=socks5://tor:9050
ports:
# Allows connectio via nginx
- "8000:8000"
- 8000:8000
tor:
image: ghcr.io/hundehausen/tor-hidden-service:latest
+4845 -869
View File
File diff suppressed because it is too large Load Diff
+109 -12
View File
@@ -1,9 +1,11 @@
import asyncio
import hashlib
import os
import json
from fastapi import HTTPException
from fastapi import HTTPException, Request
from .cashu import credit_balance, pay_out
from .cashu import credit_balance, pay_out_with_new_session
from .db import ApiKey, AsyncSession
from .models import MODELS
@@ -26,7 +28,16 @@ async def validate_bearer_key(bearer_key: str, session: AsyncSession) -> ApiKey:
Otherwise checks if the hash of the key exists.
"""
if not bearer_key:
raise HTTPException(status_code=401, detail="api-key or cashu-token required")
raise HTTPException(
status_code=401,
detail={
"error": {
"message": "API key or Cashu token required",
"type": "invalid_request_error",
"code": "missing_api_key"
}
}
)
if bearer_key.startswith("sk-"):
if exsisting_key := await session.get(ApiKey, bearer_key[3:]):
@@ -44,14 +55,69 @@ async def validate_bearer_key(bearer_key: str, session: AsyncSession) -> ApiKey:
except Exception as e:
print(f"Redemption failed: {e}")
raise HTTPException(
status_code=401, detail=f"Invalid or expired cashu key: {e}"
status_code=401,
detail={
"error": {
"message": f"Invalid or expired Cashu key: {str(e)}",
"type": "invalid_request_error",
"code": "invalid_api_key"
}
}
)
raise HTTPException(status_code=401, detail="Invalid API key")
raise HTTPException(
status_code=401,
detail={
"error": {
"message": "Invalid API key",
"type": "invalid_request_error",
"code": "invalid_api_key"
}
}
)
async def pay_for_request(key: ApiKey, session: AsyncSession) -> None:
async def pay_for_request(key: ApiKey, session: AsyncSession, request: Request | None, request_body: bytes | None = None) -> None:
if MODEL_BASED_PRICING and os.path.exists("models.json"):
if request_body:
body = json.loads(request_body)
else:
body = await request.json()
if request_model := body.get("model"):
if request_model not in [model.id for model in MODELS]:
raise HTTPException(
status_code=400,
detail={
"error": {
"message": f"Invalid model: {request_model}",
"type": "invalid_request_error",
"code": "model_not_found"
}
}
)
model = next(model for model in MODELS if model.id == request_model)
if key.balance < model.sats_pricing.max_cost * 1000:
raise HTTPException(
status_code=413,
detail={
"error": {
"message": f"This model requires a minimum balance of {model.sats_pricing.max_cost} sats",
"type": "insufficient_quota",
"code": "insufficient_balance"
}
}
)
if key.balance < COST_PER_REQUEST:
raise HTTPException(status_code=402, detail=f"Insufficient balance: {COST_PER_REQUEST} mSats required. {key.balance} available.")
raise HTTPException(
status_code=402,
detail={
"error": {
"message": f"Insufficient balance: {COST_PER_REQUEST} mSats required. {key.balance} available.",
"type": "insufficient_quota",
"code": "insufficient_balance"
}
}
)
# Charge the base cost for the request
key.balance -= COST_PER_REQUEST
@@ -77,19 +143,47 @@ async def adjust_payment_for_tokens(
"total_msats": COST_PER_REQUEST,
}
# Check if we have usage data
if "usage" not in response_data or response_data["usage"] is None:
print("No usage data in response, using base cost only")
return cost_data
# Default to configured pricing
MSATS_PER_1K_INPUT_TOKENS = COST_PER_1K_INPUT_TOKENS
MSATS_PER_1K_OUTPUT_TOKENS = COST_PER_1K_OUTPUT_TOKENS
if MODEL_BASED_PRICING and os.path.exists("models.json"):
response_model = response_data.get("model", "")
if response_model not in [model.id for model in MODELS]:
raise HTTPException(status_code=400, detail="Invalid model")
raise HTTPException(
status_code=400,
detail={
"error": {
"message": f"Invalid model in response: {response_model}",
"type": "invalid_request_error",
"code": "model_not_found"
}
}
)
model = next(model for model in MODELS if model.id == response_model)
if model.sats_pricing is None:
raise HTTPException(status_code=400, detail="Model pricing not defined")
# TODO: Rename, This is named very close to COST_PER_1K_OUTPUT_TOKENS
raise HTTPException(
status_code=400,
detail={
"error": {
"message": "Model pricing not defined",
"type": "invalid_request_error",
"code": "pricing_not_found"
}
}
)
MSATS_PER_1K_INPUT_TOKENS = model.sats_pricing.prompt * 1_000_000
MSATS_PER_1K_OUTPUT_TOKENS = model.sats_pricing.completion * 1_000_000
if not (MSATS_PER_1K_OUTPUT_TOKENS and MSATS_PER_1K_INPUT_TOKENS):
raise HTTPException(status_code=400, detail="Model pricing not defined")
# If no token pricing is configured, just return base cost
return cost_data
input_tokens = response_data.get("usage", {}).get("prompt_tokens", 0)
output_tokens = response_data.get("usage", {}).get("completion_tokens", 0)
@@ -117,6 +211,9 @@ async def adjust_payment_for_tokens(
f"Warning: Insufficient balance for token-based pricing adjustment: {key.hashed_key[:10]}..."
)
# Still proceed but log the issue - we already provided the service
# Add information about insufficient balance to cost data
cost_data["warning"] = "Insufficient balance for full token-based pricing"
cost_data["balance_shortage_msats"] = cost_difference - key.balance
else:
key.balance -= cost_difference
key.total_spent += cost_difference
@@ -131,6 +228,6 @@ async def adjust_payment_for_tokens(
session.add(key)
await session.commit()
await pay_out(session)
asyncio.create_task(pay_out_with_new_session())
return cost_data
+48 -25
View File
@@ -10,6 +10,7 @@ from .db import ApiKey, AsyncSession
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))
DEV_LN_ADDRESS = "routstr@minibits.cash"
DEVS_DONATION_RATE = 0.021 # 2.1%
WALLET = None
@@ -85,42 +86,64 @@ async def _pay_invoice_with_cashu(
return quote.amount
async def pay_out_with_new_session() -> None:
"""
Wrapper for pay_out that creates its own database session.
This prevents database connection conflicts when called as a background task.
"""
from .db import create_session
try:
async with create_session() as session:
await pay_out(session)
except Exception as e:
print(f"Error in pay_out_with_new_session: {e}")
async def pay_out(session: AsyncSession) -> None:
"""
Calculates the pay-out amount based on the spent balance, profit, and donation rate.
"""
balance = (
await session.exec(
select(func.sum(col(ApiKey.balance))).where(ApiKey.balance > 0)
)
).one()
if balance is None:
raise ValueError("No balance to pay out.")
user_balance = balance // 1000 # conversion to sats
wallet = await _initialize_wallet()
wallet_balance = wallet.available_balance
try:
balance = (
await session.exec(
select(func.sum(col(ApiKey.balance))).where(ApiKey.balance > 0)
)
).one()
if balance is None or balance == 0:
# No balance to pay out - this is OK, not an error
return
user_balance_sats = balance // 1000 # Convert msats to sats
wallet = await _initialize_wallet()
wallet_balance_sats = wallet.available_balance # Already in sats
assert wallet_balance <= user_balance, "Something went deeply wrong."
# Handle edge cases more gracefully
if wallet_balance_sats < user_balance_sats:
print(f"Warning: Wallet balance ({wallet_balance_sats} sats) is less than user balance ({user_balance_sats} sats). Skipping payout.")
return
print(f"Wallet-balance: {wallet_balance}, User-balance: {user_balance}, Revenue: {wallet_balance - user_balance}, MinPayout:{MINIMUM_PAYOUT}", flush=True)
# Why is that bad?
#assert wallet_balance <= user_balance, f"Something went deeply wrong. Wallet-balance: {wallet_balance}, User-Balance: {user_balance}"
if (revenue := wallet_balance - user_balance) <= MINIMUM_PAYOUT:
return
if (revenue := wallet_balance_sats - user_balance_sats) <= MINIMUM_PAYOUT:
# Not enough revenue yet - this is OK
return
devs_donation = int(revenue * DEVS_DONATION_RATE)
owners_draw = revenue - devs_donation
devs_donation = int(revenue * DEVS_DONATION_RATE)
owners_draw = revenue - devs_donation
await send_to_lnurl(wallet, RECEIVE_LN_ADDRESS, owners_draw * 1000) # conversion to msats for send_to_lnurl
if devs_donation > 0:
# Send payouts
print(f"Sending {owners_draw} sats to {RECEIVE_LN_ADDRESS}")
await send_to_lnurl(wallet, RECEIVE_LN_ADDRESS, owners_draw * 1000) # Convert to msats
print(f"Sending {devs_donation} sats to {DEV_LN_ADDRESS}")
await send_to_lnurl(
wallet,
"npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash",
devs_donation * 1000,
DEV_LN_ADDRESS,
devs_donation * 1000, # Convert to msats
)
except Exception as e:
# Log the error but don't crash - payouts can be retried later
print(f"Error in pay_out: {e}")
async def credit_balance(cashu_token: str, key: ApiKey, session: AsyncSession) -> int:
token_obj: Token = deserialize_token_from_string(cashu_token)
@@ -246,7 +269,7 @@ async def get_lnurl_data(lnurl: str) -> tuple[str, int, int]:
elif lnurl.lower().startswith("lnurl"):
try:
# Optional import for environments where bech32 might not be present initially
from bech32 import bech32_decode, convertbits
from bech32 import bech32_decode, convertbits # type: ignore
hrp, data = bech32_decode(lnurl)
if data is None:
+206
View File
@@ -0,0 +1,206 @@
from fastapi import APIRouter
import asyncio
import json
import websockets
import random
import string
import re
import httpx
import os
providers_router = APIRouter(prefix="/v1/providers")
def generate_subscription_id() -> str:
"""Generate a random subscription ID."""
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,
relay_url: str,
kinds: list[int] | None = None,
limit: int = 1000,
timeout: int = 30,
) -> list[dict]:
"""
Query a Nostr relay and filter for events containing a search term.
"""
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:
# 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,
}
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}")
await websocket.send(req_message)
while True:
try:
message = await asyncio.wait_for(websocket.recv(), timeout=5)
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])
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")
break
except json.JSONDecodeError:
print("Failed to decode message as JSON")
continue
await websocket.send(json.dumps(["CLOSE", sub_id]))
except Exception as e:
print(f"Query failed: {e}")
print(f"Query complete. Found {len(events)} matching events")
return events
async def get_cache() -> list[dict]:
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."""
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")
# Configure httpx to use Tor SOCKS5 proxy
async with httpx.AsyncClient(
proxies={"http://": tor_proxy, "https://": tor_proxy},
timeout=httpx.Timeout(30.0),
follow_redirects=True,
) 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"}}
@providers_router.get("/")
async def get_providers(include_json: bool = False):
npub = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s"
# Relays that support NIP-50 text search
search_relays = [
"wss://relay.nostr.band", # Known to support search
"wss://nostr.wine", # Known to support search
"wss://relay.damus.io",
"wss://nos.lol",
]
# 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}")
try:
events = await query_nostr_relay_with_search(
search_term=search_term,
relay_url=relay_url,
kinds=[1], # Text notes
limit=500,
)
# Add unique events
for event in events:
if event["id"] not in event_ids:
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
except Exception as e:
print(f"Failed to query {relay_url}: {e}")
continue
print(f"Found {len(all_events)} total unique events mentioning routstr")
providers = []
for event in all_events:
onion_urls = extract_onion_urls(event["content"])
providers.extend(onion_urls)
unique_providers = list(set(providers))
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)
if include_json:
healthy_providers.append({provider: response["json"]})
else:
healthy_providers.append(provider)
return {"providers": healthy_providers}
+2
View File
@@ -9,6 +9,7 @@ from .proxy import proxy_router
from .account import account_router
from .cashu import _initialize_wallet
from .models import MODELS, update_sats_pricing
from .discovery import providers_router
__version__ = "0.0.1"
@@ -45,6 +46,7 @@ async def info():
app.include_router(admin_router)
app.include_router(account_router)
app.include_router(providers_router)
app.include_router(proxy_router)
+21 -2
View File
@@ -20,7 +20,12 @@ class Pricing(BaseModel):
image: float
web_search: float
internal_reasoning: float
max_cost: float = 0.0 # in sats not msats
class TopProvider(BaseModel):
context_length: int | None = None
max_completion_tokens: int | None = None
is_moderated: bool | None = None
class Model(BaseModel):
id: str
@@ -30,8 +35,9 @@ class Model(BaseModel):
context_length: int
architecture: Architecture
pricing: Pricing
sats_pricing: Pricing | None
per_request_limits: dict | None
sats_pricing: Pricing | None = None
per_request_limits: dict | None = None
top_provider: TopProvider | None = None
MODELS: list[Model] = []
@@ -48,6 +54,19 @@ async def update_sats_pricing() -> None:
model.sats_pricing = Pricing(
**{k: v / sats_to_usd for k, v in model.pricing.dict().items()}
)
if model.top_provider:
if model.top_provider.context_length and model.top_provider.max_completion_tokens:
max_context_cost = model.top_provider.context_length * model.sats_pricing.prompt
max_completion_cost = model.top_provider.max_completion_tokens * model.sats_pricing.completion
model.sats_pricing.max_cost = max_context_cost + max_completion_cost
else:
p = model.sats_pricing.prompt * 1_000_000
c = model.sats_pricing.completion * 32_000
r = model.sats_pricing.request * 100_000
i = model.sats_pricing.image * 100
w = model.sats_pricing.web_search * 1000
ir = model.sats_pricing.internal_reasoning * 100
model.sats_pricing.max_cost = p + c + r + i + w + ir
except Exception as e:
print(e)
await asyncio.sleep(10)
+110 -30
View File
@@ -6,7 +6,7 @@ from fastapi.responses import Response, StreamingResponse
import httpx
import re
from router.cashu import pay_out
from router.cashu import pay_out_with_new_session
from .auth import validate_bearer_key, pay_for_request, adjust_payment_for_tokens
from .db import AsyncSession, get_session
@@ -27,7 +27,41 @@ async def proxy(
bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else ""
key = await validate_bearer_key(bearer_key, session)
await pay_for_request(key, session)
# Pre-validate JSON for requests that require it
request_body = None
if request.method in ["POST", "PUT", "PATCH"] and path.endswith("chat/completions"):
try:
request_body = await request.body()
# Try to parse JSON to validate it
if request_body:
json.loads(request_body)
except json.JSONDecodeError as e:
return Response(
content=json.dumps({
"error": {
"message": f"Invalid JSON in request body: {str(e)}",
"type": "invalid_request_error",
"code": "invalid_json"
}
}),
status_code=400,
media_type="application/json"
)
except Exception as e:
return Response(
content=json.dumps({
"error": {
"message": "Error reading request body",
"type": "invalid_request_error",
"code": "request_error"
}
}),
status_code=400,
media_type="application/json"
)
await pay_for_request(key, session, request, request_body)
# Prepare headers, removing sensitive/problematic ones
headers = dict(request.headers)
@@ -45,19 +79,35 @@ async def proxy(
path = path.replace("v1/", "")
url = f"{UPSTREAM_BASE_URL}/{path}"
client = httpx.AsyncClient(transport=httpx.AsyncHTTPTransport(retries=1))
client = httpx.AsyncClient(
transport=httpx.AsyncHTTPTransport(retries=1),
timeout=None # No timeout - requests can take as long as needed
)
try:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=request.stream(),
params=request.query_params,
),
stream=True,
)
# Use the pre-read body if available, otherwise stream
if request_body is not None:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=request_body,
params=request.query_params,
),
stream=True,
)
else:
response = await client.send(
client.build_request(
request.method,
url,
headers=headers,
content=request.stream(),
params=request.query_params,
),
stream=True,
)
# For chat completions, we need to handle token-based pricing
if path.endswith("chat/completions"):
@@ -114,17 +164,10 @@ async def proxy(
usage_data_found = True
break
except json.JSONDecodeError:
# Not valid JSON, skip
continue
if usage_data_found:
break
except Exception as e:
print(f"Error processing chunk for cost: {e}")
if not usage_data_found:
print("No usage data found in any chunks")
print(f"Error processing streaming response for cost: {e}")
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -146,7 +189,6 @@ async def proxy(
key, response_json, session
)
response_json["cost"] = cost_data
asyncio.create_task(pay_out(session))
return Response(
content=json.dumps(response_json).encode(),
status_code=response.status_code,
@@ -165,8 +207,8 @@ async def proxy(
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
background_tasks.add_task(pay_out_with_new_session)
asyncio.create_task(pay_out(session))
return StreamingResponse(
response.aiter_bytes(),
status_code=response.status_code,
@@ -176,15 +218,53 @@ async def proxy(
except httpx.RequestError as exc:
await client.aclose()
print(f"Error forwarding request to upstream: {exc}")
error_type = type(exc).__name__
error_details = str(exc)
print(
f"Error forwarding request to upstream: {error_type}: {error_details}\n"
f"Request details: method={request.method}, url={url}, headers={headers}, "
f"path={path}, query_params={dict(request.query_params)}"
)
# Provide more specific error messages based on the error type
if isinstance(exc, httpx.ConnectError):
error_message = "Unable to connect to upstream service"
elif isinstance(exc, httpx.TimeoutException):
error_message = "Upstream service request timed out"
elif isinstance(exc, httpx.NetworkError):
error_message = "Network error while connecting to upstream service"
else:
error_message = f"Error connecting to upstream service: {error_type}"
return Response(
content=f"Error connecting to upstream service: {exc}",
content=json.dumps({
"error": {
"message": error_message,
"type": "upstream_error",
"code": 502
}
}),
status_code=502,
media_type="application/json"
)
except Exception as exc:
await client.aclose()
print(f"Unexpected error: {exc}")
return Response(
content=f"Unexpected server error: {exc}",
status_code=500,
import traceback
tb = traceback.format_exc()
print(
f"Unexpected error: {exc}\n"
f"Request details: method={request.method}, url={url}, headers={headers}, "
f"path={path}, query_params={dict(request.query_params)}\n"
f"Traceback:\n{tb}"
)
return Response(
content=json.dumps({
"error": {
"message": "An unexpected server error occurred",
"type": "internal_error",
"code": 500
}
}),
status_code=500,
media_type="application/json"
)