mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge branch 'main' into sixty-nuts-migration
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
@@ -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)"
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user