Merge branch 'main' into sixty-nuts-migration

This commit is contained in:
Shroominic
2025-06-03 15:48:33 +02:00
12 changed files with 199 additions and 47 deletions
+8
View File
@@ -9,6 +9,9 @@ UPSTREAM_API_KEY="sk-21212121212121212121212121212121"
# Lightning address used to receive funds
RECEIVE_LN_ADDRESS="shroominic@walletofsatoshi.com"
# 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"
@@ -17,6 +20,11 @@ COST_PER_1K_OUTPUT_TOKENS = "0"
# If set to true, make sure model pricings are defined in models.json
MODEL_BASED_PRICING = "false"
# 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"
# password used to log into admin interface
ADMIN_PASSWORD="XXX"
+4 -3
View File
@@ -5,7 +5,8 @@ wallet.sqlite3
# Development
.notes
.keys.db
.wallet.sqlite3
.models.json
.*keys.db
.*wallet.sqlite3
.*models.json
compose.override.yml
+29 -8
View File
@@ -1,4 +1,5 @@
import os
from datetime import datetime, timezone
from fastapi import APIRouter, Request
from fastapi.responses import HTMLResponse
@@ -92,14 +93,30 @@ async def dashboard(request: Request) -> str:
async with create_session() as session:
result = await session.exec(select(ApiKey))
api_keys = result.all()
api_keys_table_rows = "".join(
f"<tr><td>{key.hashed_key}</td><td>{key.balance}</td><td>{key.refund_address}</td><td>{key.total_spent}</td><td>{key.total_requests}</td></tr>"
for key in api_keys
)
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>"
)
api_keys_table_rows = "".join(api_keys_table_rows)
# Calculate the total balance of all API keys
total_user_balance = int(sum(key.balance / 1000 for key in api_keys))
# Fetch balance from cashu
async with Wallet(nsec=NSEC, mint_urls=[MINT]) as wallet:
current_balance = (await wallet.fetch_wallet_state()).balance
owner_balance = current_balance - total_user_balance
return f"""<!DOCTYPE html>
<html>
@@ -118,16 +135,20 @@ async def dashboard(request: Request) -> str:
</head>
<body>
<h1>Admin Dashboard</h1>
<h2>Current Cashu Balance (including user balances)</h2>
<p>{current_balance} sats</p>
<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>Refund Address</th>
<th>Total Spent(mSats)</th>
<th>Total Spent (mSats)</th>
<th>Total Requests</th>
<th>Refund Address</th>
<th>Refund Time</th>
</tr>
{api_keys_table_rows}
</table>
+30 -7
View File
@@ -2,6 +2,8 @@ import asyncio
import hashlib
import os
import json
from typing import Optional
from fastapi import HTTPException, Request
@@ -21,7 +23,12 @@ COST_PER_1K_OUTPUT_TOKENS = (
MODEL_BASED_PRICING = os.environ.get("MODEL_BASED_PRICING", "false").lower() == "true"
async def validate_bearer_key(bearer_key: str, session: AsyncSession) -> ApiKey:
async def validate_bearer_key(
bearer_key: str,
session: AsyncSession,
refund_address: Optional[str] = None,
key_expiry_time: Optional[int] = None,
) -> ApiKey:
"""
Validates the provided API key using SQLModel.
If it's a cashu key, it redeems it and stores its hash and balance.
@@ -40,16 +47,32 @@ async def validate_bearer_key(bearer_key: str, session: AsyncSession) -> ApiKey:
)
if bearer_key.startswith("sk-"):
if exsisting_key := await session.get(ApiKey, bearer_key[3:]):
return exsisting_key
if existing_key := await session.get(ApiKey, bearer_key[3:]):
existing_key.key_expiry_time, existing_key.refund_address = (
key_expiry_time,
refund_address,
)
return existing_key
if bearer_key.startswith("cashu"):
try:
hashed_key = hashlib.sha256(bearer_key.encode()).hexdigest()
if exsisting_key := await session.get(ApiKey, hashed_key):
return exsisting_key
new_key = ApiKey(hashed_key=hashed_key, balance=0)
await credit_balance(bearer_key, new_key, session)
if existing_key := await session.get(ApiKey, hashed_key):
existing_key.key_expiry_time, existing_key.refund_address = (
key_expiry_time,
refund_address,
)
return existing_key
new_key = ApiKey(
hashed_key=hashed_key,
balance=0,
refund_address=refund_address,
key_expiry_time=key_expiry_time,
)
await credit_balance(
bearer_key, new_key, session
) # TODO: see cashu.py "_initialize_wallet"
await session.refresh(new_key)
return new_key
except Exception as e:
+44 -1
View File
@@ -1,15 +1,20 @@
import os
import httpx
import asyncio
import time
from sixty_nuts import Wallet
from sqlmodel import select, func, col
from .db import ApiKey, AsyncSession, get_session
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))
REFUND_PROCESSING_INTERVAL = int(os.environ.get("REFUND_PROCESSING_INTERVAL", 3600))
DEV_LN_ADDRESS = "routstr@minibits.cash"
DEVS_DONATION_RATE = 0.021 # 2.1%
DEVS_DONATION_RATE = float(os.environ.get("DEVS_DONATION_RATE", 0.021)) # 2.1%
NSEC = os.environ["NSEC"] # Nostr private key for the wallet
@@ -71,6 +76,44 @@ async def credit_balance(cashu_token: str, key: ApiKey, session: AsyncSession) -
return amount
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.
"""
raise Exception("TODO migrate to sixty-nuts")
# Setting REFUND_PROCESSING_INTERVAL to 0 disables it
if REFUND_PROCESSING_INTERVAL == 0:
print("Automatic refund processing is disabled.")
return
while True:
try:
async for session in get_session():
result = await session.exec(select(ApiKey))
keys = result.all()
current_time = int(time.time())
for key in keys:
if (
key.balance > 0
and key.refund_address
and key.key_expiry_time
and key.key_expiry_time < current_time
):
print(
f" DEBUG Refunding key {key.hashed_key[:3] + '[...]' + key.hashed_key[-3:]}, Current Time: {current_time}, Expirary Time: {key.key_expiry_time}",
flush=True,
)
await refund_balance(key.balance, key, session)
# Sleep for the specified interval before checking again
await asyncio.sleep(REFUND_PROCESSING_INTERVAL)
except Exception as e:
print(f"Error during refund check: {e}")
async def refund_balance(amount: int, key: ApiKey, session: AsyncSession) -> int:
async with Wallet(nsec=NSEC, mint_urls=[MINT]) as wallet:
if key.balance < amount:
+4
View File
@@ -21,6 +21,10 @@ class ApiKey(SQLModel, table=True): # type: ignore
default=None,
description="Lightning address to refund remaining balance after key expires",
)
key_expiry_time: int | None = Field(
default=None,
description="Unix-timestamp after which the cashu-token's balance gets refunded to the refund_address",
)
total_spent: int = Field(
default=0, description="Total spent in millisatoshis (msats)"
)
+3
View File
@@ -8,8 +8,10 @@ from .admin import admin_router
from .proxy import proxy_router
from .account import account_router
from .models import MODELS, update_sats_pricing
from .cashu import check_for_refunds
from .discovery import providers_router
__version__ = "0.0.1"
app = FastAPI(
@@ -53,3 +55,4 @@ app.include_router(proxy_router)
async def startup_event():
await init_db()
asyncio.create_task(update_sats_pricing())
asyncio.create_task(check_for_refunds())
+35 -3
View File
@@ -24,8 +24,32 @@ async def proxy(
):
auth = request.headers.get("Authorization", "")
bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else ""
refund_address = request.headers.get("Refund-LNURL", None)
key_expiry_time = request.headers.get("Key-Expiry-Time", None)
key = await validate_bearer_key(bearer_key, session)
# Validate key_expiry_time header
if key_expiry_time:
try:
key_expiry_time = int(key_expiry_time) # type: ignore
except ValueError:
return Response(
content="Invalid Key-Expiry-Time: must be a valid Unix timestamp",
status_code=400,
)
if not refund_address:
return Response(
content="Error: Refund-LNURL header required when using Key-Expiry-Time",
status_code=400,
)
else:
key_expiry_time = None
key = await validate_bearer_key(
bearer_key,
session,
refund_address,
key_expiry_time, # type: ignore
)
# Pre-validate JSON for requests that require it
request_body = None
@@ -70,6 +94,8 @@ async def proxy(
headers = dict(request.headers)
headers.pop("host", None)
headers.pop("content-length", None)
headers.pop("refund-lnurl", None)
headers.pop("key-expiry-time", None)
if UPSTREAM_API_KEY:
headers["Authorization"] = f"Bearer {UPSTREAM_API_KEY}"
@@ -175,7 +201,6 @@ async def proxy(
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
return StreamingResponse(
stream_with_cost(),
status_code=response.status_code,
@@ -192,10 +217,17 @@ async def proxy(
key, response_json, session
)
response_json["cost"] = cost_data
response_headers = dict(response.headers)
# Remove Transfer-Encoding header to avoid conflict with Content-Length header in common nginx setups
if "transfer-encoding" in response_headers:
del response_headers["transfer-encoding"]
return Response(
content=json.dumps(response_json).encode(),
status_code=response.status_code,
headers=dict(response.headers),
headers=response_headers,
media_type="application/json",
)
except json.JSONDecodeError as e:
+1 -1
View File
@@ -119,7 +119,7 @@ def test_client() -> TestClient:
with patch("router.models.update_sats_pricing") as mock_update:
mock_update.return_value = None
return TestClient(app)
yield TestClient(app)
@pytest_asyncio.fixture
+10 -14
View File
@@ -7,20 +7,16 @@ from unittest.mock import patch
async def test_root_endpoint(async_client: AsyncClient):
"""Test the root endpoint returns expected information."""
# Mock the environment variables for this specific test
with patch("os.environ.get") as mock_env_get:
def env_side_effect(key, default=None):
env_map = {
"NAME": "TestRoutstrNode",
"DESCRIPTION": "Test Node",
"NPUB": "npub1test",
"MINT": "https://test.mint.com",
"HTTP_URL": "http://test.example.com",
"ONION_URL": "http://test.onion",
}
return env_map.get(key, default)
mock_env_get.side_effect = env_side_effect
env_vars = {
"NAME": "TestRoutstrNode",
"DESCRIPTION": "Test Node",
"NPUB": "npub1test",
"MINT": "https://test.mint.com",
"HTTP_URL": "http://test.example.com",
"ONION_URL": "http://test.onion",
}
with patch.dict("os.environ", env_vars, clear=False):
response = await async_client.get("/")
assert response.status_code == 200
+30 -9
View File
@@ -67,14 +67,21 @@ async def test_update_sats_pricing_calculation(sample_model: Model):
assert sample_model.sats_pricing is not None
# Verify calculations (prices in USD / sats_to_usd)
assert sample_model.sats_pricing.prompt == 0.01 / 0.0001 # 100 sats
assert sample_model.sats_pricing.completion == 0.02 / 0.0001 # 200 sats
assert sample_model.sats_pricing.request == 0.001 / 0.0001 # 10 sats
assert sample_model.sats_pricing.prompt == pytest.approx(0.01 / 0.0001) # 100 sats
assert sample_model.sats_pricing.completion == pytest.approx(0.02 / 0.0001) # 200 sats
assert sample_model.sats_pricing.request == pytest.approx(0.001 / 0.0001) # 10 sats
# Verify max_cost calculation for model with top_provider
expected_max_context = 4096 * sample_model.sats_pricing.prompt
expected_max_completion = 2048 * sample_model.sats_pricing.completion
assert sample_model.sats_pricing.max_cost == expected_max_context + expected_max_completion
assert sample_model.sats_pricing.max_cost == pytest.approx(expected_max_context + expected_max_completion)
# Cancel and await the task
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
except asyncio.CancelledError:
pass
@@ -140,7 +147,14 @@ async def test_update_sats_pricing_without_top_provider():
ir = model_without_top.sats_pricing.internal_reasoning * 100
expected_max = p + c + r + i + w + ir
assert model_without_top.sats_pricing.max_cost == expected_max
assert model_without_top.sats_pricing.max_cost == pytest.approx(expected_max)
# Cancel and await the task
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
except asyncio.CancelledError:
pass
@@ -158,11 +172,11 @@ async def test_update_sats_pricing_handles_errors():
error_printed = False
original_print = print
def mock_print(msg):
def mock_print(*args, **kwargs):
nonlocal error_printed
if isinstance(msg, Exception) and str(msg) == "API Error":
if args and isinstance(args[0], Exception) and str(args[0]) == "API Error":
error_printed = True
original_print(msg)
original_print(*args, **kwargs)
with patch("builtins.print", side_effect=mock_print):
sleep_called = asyncio.Event()
@@ -179,6 +193,13 @@ async def test_update_sats_pricing_handles_errors():
# Verify error was printed
assert error_printed
# Cancel and await the task
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
except asyncio.CancelledError:
pass
@@ -197,4 +218,4 @@ def test_model_serialization(sample_model: Model):
# Test deserialization
new_model = Model(**model_dict)
assert new_model.id == sample_model.id
assert new_model.pricing.prompt == sample_model.pricing.prompt
assert new_model.pricing.prompt == pytest.approx(sample_model.pricing.prompt)
+1 -1
View File
@@ -181,7 +181,7 @@ async def test_proxy_streaming_response(
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"}
mock_response.aiter_bytes = mock_aiter_bytes
mock_response.aiter_bytes = lambda: mock_aiter_bytes()
mock_response.aclose = AsyncMock()
mock_client.send = AsyncMock(return_value=mock_response)