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

- - + + + {api_keys_table_rows}
Hashed Key Balance (mSats)Refund AddressTotal Spent(mSats)Total Spent (mSats) Total RequestsRefund AddressRefund Time
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)