diff --git a/.env.example b/.env.example
index 9e402b08..3dcf8ec9 100644
--- a/.env.example
+++ b/.env.example
@@ -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"
diff --git a/.gitignore b/.gitignore
index 359006b7..b25f148c 100644
--- a/.gitignore
+++ b/.gitignore
@@ -5,7 +5,8 @@ wallet.sqlite3
# Development
.notes
-.keys.db
-.wallet.sqlite3
-.models.json
+.*keys.db
+.*wallet.sqlite3
+.*models.json
+
compose.override.yml
diff --git a/router/admin.py b/router/admin.py
index 5b2ff9a8..e8400237 100644
--- a/router/admin.py
+++ b/router/admin.py
@@ -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"
| {key.hashed_key} | {key.balance} | {key.refund_address} | {key.total_spent} | {key.total_requests} |
"
- 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"| {key.hashed_key} | {key.balance} | {key.total_spent} | {key.total_requests} | {key.refund_address} | {'{} ({} UTC)'.format(key.key_expiry_time, expiry_time_human_readable) if key.key_expiry_time else key.key_expiry_time} |
"
+ )
+
+ 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"""
@@ -118,16 +135,20 @@ async def dashboard(request: Request) -> str:
Admin Dashboard
- Current Cashu Balance (including user balances)
- {current_balance} sats
+ Current Cashu Balance
+ Your Balance: {owner_balance} sats
+ The balance is calculated by subtracting the combined user balance from the total Cashu wallet balance.
+ Total Cashu Balance: {current_balance} sats
+ User Balance: {total_user_balance} sats
User's API Keys
| Hashed Key |
Balance (mSats) |
- Refund Address |
- Total Spent(mSats) |
+ Total Spent (mSats) |
Total Requests |
+ Refund Address |
+ Refund Time |
{api_keys_table_rows}
diff --git a/router/auth.py b/router/auth.py
index aad8c6cc..978478e2 100644
--- a/router/auth.py
+++ b/router/auth.py
@@ -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:
diff --git a/router/cashu.py b/router/cashu.py
index 63f4bb17..713b3fe7 100644
--- a/router/cashu.py
+++ b/router/cashu.py
@@ -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:
diff --git a/router/db.py b/router/db.py
index dec53645..10e45d8b 100644
--- a/router/db.py
+++ b/router/db.py
@@ -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)"
)
diff --git a/router/main.py b/router/main.py
index 21bd327b..d266fd6f 100644
--- a/router/main.py
+++ b/router/main.py
@@ -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())
diff --git a/router/proxy.py b/router/proxy.py
index 2b16b3be..0f358d61 100644
--- a/router/proxy.py
+++ b/router/proxy.py
@@ -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:
diff --git a/tests/conftest.py b/tests/conftest.py
index ed4ff014..723d35d4 100644
--- a/tests/conftest.py
+++ b/tests/conftest.py
@@ -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
diff --git a/tests/test_main.py b/tests/test_main.py
index 0a66632a..cee18b37 100644
--- a/tests/test_main.py
+++ b/tests/test_main.py
@@ -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
diff --git a/tests/test_models.py b/tests/test_models.py
index d28a2f2f..c4a8494e 100644
--- a/tests/test_models.py
+++ b/tests/test_models.py
@@ -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
\ No newline at end of file
+ assert new_model.pricing.prompt == pytest.approx(sample_model.pricing.prompt)
\ No newline at end of file
diff --git a/tests/test_proxy.py b/tests/test_proxy.py
index d44f86ed..c8301299 100644
--- a/tests/test_proxy.py
+++ b/tests/test_proxy.py
@@ -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)