From b70b94b9b4e2e3f4e065a4a02c204eeb0ac6efe0 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 3 Jan 2026 22:12:23 +0100 Subject: [PATCH 01/19] feat: add admin password warning log --- routstr/core/main.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/routstr/core/main.py b/routstr/core/main.py index d36b4447..44ed56a2 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -62,6 +62,11 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: async with create_session() as session: s = await SettingsService.initialize(session) + if not s.admin_password: + logger.info( + f"Admin password is not set. Visit {s.http_url or 'http://localhost:8000'}/admin to set the password." + ) + # Apply app metadata from settings try: app.title = s.name From 6d780ef96dbb205134dde8963c67a08924805849 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 3 Jan 2026 22:12:55 +0100 Subject: [PATCH 02/19] feat: disable provider discovery by default --- routstr/core/settings.py | 2 +- routstr/discovery.py | 5 ++++- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 302df24d..6cbec561 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -59,7 +59,7 @@ class Settings(BaseSettings): cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL") providers_refresh_interval_seconds: int = Field( - default=300, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" + default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" ) pricing_refresh_interval_seconds: int = Field( default=120, env="PRICING_REFRESH_INTERVAL_SECONDS" diff --git a/routstr/discovery.py b/routstr/discovery.py index 03a82703..94375b63 100644 --- a/routstr/discovery.py +++ b/routstr/discovery.py @@ -6,7 +6,7 @@ from typing import Any import httpx import websockets -from fastapi import APIRouter +from fastapi import APIRouter, HTTPException from .core.logging import get_logger from .core.settings import settings @@ -389,6 +389,9 @@ async def get_providers( Return cached providers. If include_json, return provider+health; otherwise provider only. Optional filter by pubkey. """ + if settings.providers_refresh_interval_seconds == 0: + raise HTTPException(status_code=404, detail="Provider discovery is disabled") + cache = await get_cache() if not cache: await refresh_providers_cache(pubkey=pubkey) From 9e9bc5bff8efd280c1fb4be5672e75515166b855 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 3 Jan 2026 22:49:46 +0100 Subject: [PATCH 03/19] fix pytests --- tests/integration/test_performance_load.py | 61 ++++++++++--------- tests/integration/test_provider_management.py | 7 +++ 2 files changed, 40 insertions(+), 28 deletions(-) diff --git a/tests/integration/test_performance_load.py b/tests/integration/test_performance_load.py index e5e5c3ce..bc74a60f 100644 --- a/tests/integration/test_performance_load.py +++ b/tests/integration/test_performance_load.py @@ -9,6 +9,7 @@ import gc import statistics import time from typing import Any, Dict, List +from unittest.mock import patch import psutil import pytest @@ -105,42 +106,46 @@ class TestPerformanceBaseline: ("GET", "/v1/wallet/info", authenticated_client, None), ] - # Warm up - for _ in range(10): - await integration_client.get("/") + # Enable provider discovery for this test + with patch( + "routstr.core.settings.settings.providers_refresh_interval_seconds", 300 + ): + # Warm up + for _ in range(10): + await integration_client.get("/") - # Test each endpoint - for method, path, client, data in endpoints: - response_times = [] + # Test each endpoint + for method, path, client, data in endpoints: + response_times = [] - for i in range(100): - start = time.time() + for i in range(100): + start = time.time() - if method == "GET": - response = await client.get(path) - else: - response = await client.post(path, json=data) + if method == "GET": + response = await client.get(path) + else: + response = await client.post(path, json=data) - duration = time.time() - start - response_times.append(duration * 1000) # Convert to ms + duration = time.time() - start + response_times.append(duration * 1000) # Convert to ms - assert response.status_code in [200, 201] + assert response.status_code in [200, 201] - if i % 10 == 0: - metrics.record_system_metrics() + if i % 10 == 0: + metrics.record_system_metrics() - # Verify 95th percentile < 500ms - p95 = sorted(response_times)[int(len(response_times) * 0.95)] - assert p95 < 500, ( - f"{method} {path} p95 response time {p95}ms exceeds 500ms limit" - ) + # Verify 95th percentile < 500ms + p95 = sorted(response_times)[int(len(response_times) * 0.95)] + assert p95 < 500, ( + f"{method} {path} p95 response time {p95}ms exceeds 500ms limit" + ) - print(f"\n{method} {path}:") - print(f" Mean: {statistics.mean(response_times):.2f}ms") - print(f" P95: {p95:.2f}ms") - print( - f" P99: {sorted(response_times)[int(len(response_times) * 0.99)]:.2f}ms" - ) + print(f"\n{method} {path}:") + print(f" Mean: {statistics.mean(response_times):.2f}ms") + print(f" P95: {p95:.2f}ms") + print( + f" P99: {sorted(response_times)[int(len(response_times) * 0.99)]:.2f}ms" + ) @pytest.mark.integration diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index 788ad768..050ae4dd 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -19,6 +19,13 @@ def _clear_providers_cache() -> None: _PROVIDERS_CACHE.clear() +@pytest.fixture(autouse=True) +def _enable_provider_discovery() -> None: + """Enable provider discovery for all tests in this module""" + with patch("routstr.core.settings.settings.providers_refresh_interval_seconds", 300): + yield + + @pytest.mark.integration @pytest.mark.asyncio async def test_providers_endpoint_default_response( From 7d829af681a0bcd0b150562aa0198e161e05e1f8 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Mon, 5 Jan 2026 00:06:53 +0100 Subject: [PATCH 04/19] fix types --- tests/integration/test_provider_management.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index 050ae4dd..b4f2ad6b 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -3,7 +3,7 @@ Integration tests for provider management functionality. Tests GET /v1/providers/ endpoint for listing and managing providers. """ -from typing import Any +from typing import Any, Generator from unittest.mock import patch import pytest @@ -20,9 +20,11 @@ def _clear_providers_cache() -> None: @pytest.fixture(autouse=True) -def _enable_provider_discovery() -> None: +def _enable_provider_discovery() -> Generator[None, Any, Any]: """Enable provider discovery for all tests in this module""" - with patch("routstr.core.settings.settings.providers_refresh_interval_seconds", 300): + with patch( + "routstr.core.settings.settings.providers_refresh_interval_seconds", 300 + ): yield From 50eabafa57e776b8f82c3d937a54203f377085b1 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Mon, 5 Jan 2026 12:12:04 +0100 Subject: [PATCH 05/19] remove performance tests due to unpredictable behaviour --- tests/integration/test_example.py | 24 --------- tests/integration/test_provider_management.py | 42 +--------------- tests/integration/test_proxy_get_endpoints.py | 33 ------------- .../integration/test_proxy_post_endpoints.py | 49 ------------------- tests/integration/test_wallet_information.py | 28 +---------- tests/integration/test_wallet_refund.py | 38 -------------- 6 files changed, 2 insertions(+), 212 deletions(-) diff --git a/tests/integration/test_example.py b/tests/integration/test_example.py index 6657517c..030f7bcb 100644 --- a/tests/integration/test_example.py +++ b/tests/integration/test_example.py @@ -10,7 +10,6 @@ from httpx import AsyncClient from .utils import ( CashuTokenGenerator, - PerformanceValidator, ResponseValidator, ) @@ -159,30 +158,7 @@ async def test_error_handling( assert response.status_code == 401 -@pytest.mark.integration -@pytest.mark.asyncio -async def test_performance_requirements(integration_client: AsyncClient) -> None: - """Test that endpoints meet performance requirements""" - validator = PerformanceValidator() - - # Test info endpoint performance - for i in range(50): - start = validator.start_timing("info_endpoint") - response = await integration_client.get("/") - validator.end_timing("info_endpoint", start) - assert response.status_code == 200 - - # Validate 95th percentile is under 500ms - result = validator.validate_response_time( - "info_endpoint", max_duration=0.5, percentile=0.95 - ) - - assert result["valid"], ( - f"Performance requirement failed: " - f"95th percentile was {result['percentile_time']:.3f}s " - f"(required < {result['max_allowed']}s)" - ) @pytest.mark.integration diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index 788ad768..cb0417dc 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -11,7 +11,7 @@ from httpx import AsyncClient from routstr.discovery import _PROVIDERS_CACHE -from .utils import PerformanceValidator, ResponseValidator +from .utils import ResponseValidator @pytest.fixture(autouse=True) @@ -518,46 +518,6 @@ async def test_providers_endpoint_response_format( assert isinstance(data_json["providers"], list) -@pytest.mark.integration -@pytest.mark.asyncio -async def test_providers_endpoint_performance(integration_client: AsyncClient) -> None: - """Test providers endpoint meets performance requirements""" - - # Mock quick responses to avoid network delays - mock_events: list[dict[str, Any]] = [ - { - "id": f"event{i}", - "content": f"Provider: http://provider{i}.onion", - "created_at": 1234567890 + i, - } - for i in range(5) - ] - - validator = PerformanceValidator() - - with patch( - "routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events - ): - with patch("routstr.discovery.fetch_provider_health") as mock_fetch: - mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}} - - # Test multiple requests - for i in range(10): - start = validator.start_timing("providers_endpoint") - response = await integration_client.get("/v1/providers/") - validator.end_timing("providers_endpoint", start) - - assert response.status_code == 200 - - # Validate performance (should be fast with mocked dependencies) - perf_result = validator.validate_response_time( - "providers_endpoint", - max_duration=2.0, # Allow more time since it involves multiple operations - percentile=0.95, - ) - assert perf_result["valid"], f"Performance requirement failed: {perf_result}" - - @pytest.mark.integration @pytest.mark.asyncio async def test_providers_endpoint_concurrent_requests( diff --git a/tests/integration/test_proxy_get_endpoints.py b/tests/integration/test_proxy_get_endpoints.py index ff745d8d..9621e57d 100644 --- a/tests/integration/test_proxy_get_endpoints.py +++ b/tests/integration/test_proxy_get_endpoints.py @@ -18,7 +18,6 @@ from routstr.core.db import ApiKey from .utils import ( ConcurrencyTester, - PerformanceValidator, ) @@ -551,39 +550,7 @@ async def test_proxy_get_concurrent_requests( assert response.status_code == 200 -@pytest.mark.integration -@pytest.mark.asyncio -async def test_proxy_get_performance_requirements( - integration_client: AsyncClient, authenticated_client: AsyncClient -) -> None: - """Test that GET proxy requests meet performance requirements""" - validator = PerformanceValidator() - - with patch("httpx.AsyncClient.request") as mock_request: - mock_response = AsyncMock() - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.json = MagicMock(return_value={"performance": "test"}) - mock_response.text = '{"performance": "test"}' - mock_response.iter_bytes = AsyncMock(return_value=[b'{"performance": "test"}']) - mock_request.return_value = mock_response - - # Test multiple requests for performance measurement - for i in range(20): - start = validator.start_timing("proxy_get") - response = await authenticated_client.get(f"/v1/perf-test-{i}") - validator.end_timing("proxy_get", start) - - assert response.status_code == 200 - - # Validate performance requirements - perf_result = validator.validate_response_time( - "proxy_get", - max_duration=1.0, # Should complete within 1 second - percentile=0.95, - ) - assert perf_result["valid"], f"Performance requirement failed: {perf_result}" @pytest.mark.integration diff --git a/tests/integration/test_proxy_post_endpoints.py b/tests/integration/test_proxy_post_endpoints.py index a6bc9d9f..8d5fab36 100644 --- a/tests/integration/test_proxy_post_endpoints.py +++ b/tests/integration/test_proxy_post_endpoints.py @@ -15,7 +15,6 @@ from httpx import ASGITransport, AsyncClient from .utils import ( ConcurrencyTester, - PerformanceValidator, ) @@ -290,55 +289,7 @@ async def test_proxy_post_unauthorized_access(integration_client: AsyncClient) - assert response.status_code in [400, 401] -@pytest.mark.integration -@pytest.mark.asyncio -async def test_proxy_post_performance( - integration_client: AsyncClient, authenticated_client: AsyncClient -) -> None: - """Test POST endpoint performance requirements""" - test_payload = { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": "Performance test"}], - } - - validator = PerformanceValidator() - - with patch("httpx.AsyncClient.send") as mock_send: - # Mock fast responses - async def mock_iter_bytes(*args: Any, **kwargs: Any) -> Any: - yield b'{"choices": [{"message": {"content": "Fast"}}], "usage": {"total_tokens": 5}}' - - mock_response = AsyncMock() - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - response_data = { - "choices": [{"message": {"content": "Fast"}}], - "usage": {"total_tokens": 5}, - } - mock_response.text = json.dumps(response_data) - mock_response.json = AsyncMock(return_value=response_data) - mock_response.iter_bytes = mock_iter_bytes - mock_response.aiter_bytes = mock_iter_bytes - mock_send.return_value = mock_response - - # Run multiple requests for performance measurement - for i in range(20): - start = validator.start_timing("proxy_post") - response = await authenticated_client.post( - "/v1/chat/completions", json=test_payload - ) - validator.end_timing("proxy_post", start) - - assert response.status_code == 200 - - # Validate performance - perf_result = validator.validate_response_time( - "proxy_post", - max_duration=1.5, # Allow slightly more time for POST - percentile=0.95, - ) - assert perf_result["valid"], f"Performance requirement failed: {perf_result}" @pytest.mark.integration diff --git a/tests/integration/test_wallet_information.py b/tests/integration/test_wallet_information.py index 78dacf63..3a11f4ef 100644 --- a/tests/integration/test_wallet_information.py +++ b/tests/integration/test_wallet_information.py @@ -3,7 +3,7 @@ Integration tests for wallet information retrieval endpoints. Tests GET /v1/wallet/ and GET /v1/wallet/info endpoints with various scenarios. """ -import time + from datetime import datetime, timedelta from typing import Any @@ -367,30 +367,4 @@ async def test_wallet_info_with_special_characters_in_headers( # Note: Current implementation doesn't return refund_address in response -@pytest.mark.integration -@pytest.mark.asyncio -@pytest.mark.slow -async def test_wallet_endpoints_performance(authenticated_client: AsyncClient) -> None: - """Test wallet endpoints meet performance requirements""" - # Warm up - await authenticated_client.get("/v1/wallet/") - - # Measure response times - response_times = [] - - for _ in range(50): - start_time = time.time() - response = await authenticated_client.get("/v1/wallet/") - end_time = time.time() - - assert response.status_code == 200 - response_times.append(end_time - start_time) - - # Calculate statistics - avg_time = sum(response_times) / len(response_times) - max_time = max(response_times) - - # Performance assertions - assert avg_time < 0.1 # Average should be under 100ms - assert max_time < 0.5 # No request should take more than 500ms diff --git a/tests/integration/test_wallet_refund.py b/tests/integration/test_wallet_refund.py index c2461cae..9770ec7a 100644 --- a/tests/integration/test_wallet_refund.py +++ b/tests/integration/test_wallet_refund.py @@ -537,42 +537,4 @@ async def test_refund_with_expired_key( assert response.json()["recipient"] == "expired@ln.address" -@pytest.mark.integration -@pytest.mark.asyncio -@pytest.mark.slow -async def test_refund_performance( - integration_client: AsyncClient, testmint_wallet: Any -) -> None: - """Test refund endpoint performance""" - import time - - # Create multiple API keys - api_keys = [] - for i in range(10): - token = await testmint_wallet.mint_tokens(100 + i) - # Use cashu token as Bearer auth to create API key - integration_client.headers["Authorization"] = f"Bearer {token}" - response = await integration_client.get("/v1/wallet/info") - assert response.status_code == 200 - api_keys.append(response.json()["api_key"]) - - # Measure refund times - refund_times = [] - - for api_key in api_keys: - integration_client.headers["Authorization"] = f"Bearer {api_key}" - - start_time = time.time() - response = await integration_client.post("/v1/wallet/refund") - end_time = time.time() - - assert response.status_code == 200 - refund_times.append(end_time - start_time) - - # Performance assertions - avg_time = sum(refund_times) / len(refund_times) - max_time = max(refund_times) - - assert avg_time < 0.5 # Average under 500ms - assert max_time < 1.0 # No refund takes more than 1 second From 1e21dce7350fb0c65c1f2dd66480f83905b17c9b Mon Sep 17 00:00:00 2001 From: Shroominic Date: Tue, 6 Jan 2026 18:23:01 +0100 Subject: [PATCH 06/19] change to warning log --- routstr/core/main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/routstr/core/main.py b/routstr/core/main.py index 44ed56a2..80c03403 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -63,7 +63,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: s = await SettingsService.initialize(session) if not s.admin_password: - logger.info( + logger.warning( f"Admin password is not set. Visit {s.http_url or 'http://localhost:8000'}/admin to set the password." ) From c4cc09d61eb4b7bc6e4179bf84df3502140b4d3f Mon Sep 17 00:00:00 2001 From: Shroominic Date: Fri, 9 Jan 2026 16:50:10 +0800 Subject: [PATCH 07/19] fix provider balance not displaying when 0 --- ui/app/providers/page.tsx | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/ui/app/providers/page.tsx b/ui/app/providers/page.tsx index 4624f26a..d5d8c522 100644 --- a/ui/app/providers/page.tsx +++ b/ui/app/providers/page.tsx @@ -208,7 +208,11 @@ function ProviderBalance({ return ; } - if (error || !balanceData?.ok || !balanceData.balance_data) { + if ( + error || + !balanceData?.ok || + (balanceData.balance_data === undefined || balanceData.balance_data === null) + ) { return null; } From ca7e8bec71578adccc2ae3590de9f43c78f8296b Mon Sep 17 00:00:00 2001 From: Shroominic Date: Fri, 9 Jan 2026 16:53:25 +0800 Subject: [PATCH 08/19] fmt --- ui/app/providers/page.tsx | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/ui/app/providers/page.tsx b/ui/app/providers/page.tsx index d5d8c522..dc2b6e3e 100644 --- a/ui/app/providers/page.tsx +++ b/ui/app/providers/page.tsx @@ -211,7 +211,8 @@ function ProviderBalance({ if ( error || !balanceData?.ok || - (balanceData.balance_data === undefined || balanceData.balance_data === null) + balanceData.balance_data === undefined || + balanceData.balance_data === null ) { return null; } From 493b4f0f1f891203f9c1c83d8ce578786a66757a Mon Sep 17 00:00:00 2001 From: Shroominic Date: Fri, 9 Jan 2026 17:32:35 +0800 Subject: [PATCH 09/19] urgent reserved balance fix --- routstr/balance.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/routstr/balance.py b/routstr/balance.py index 62358761..6697a5db 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -154,7 +154,8 @@ async def refund_wallet_endpoint( return cached key: ApiKey = await validate_bearer_key(bearer_value, session) - remaining_balance_msats: int = key.balance + + remaining_balance_msats: int = key.total_balance if key.refund_currency == "sat": remaining_balance = remaining_balance_msats // 1000 From 39657ed64f7985660c7ec4c51d884053bb4fc52c Mon Sep 17 00:00:00 2001 From: Shroominic Date: Fri, 9 Jan 2026 17:36:04 +0800 Subject: [PATCH 10/19] reset reserved balance on startup --- routstr/core/db.py | 10 +++++++++- routstr/core/main.py | 4 ++++ routstr/core/settings.py | 3 +++ 3 files changed, 16 insertions(+), 1 deletion(-) diff --git a/routstr/core/db.py b/routstr/core/db.py index c56163d0..4c236d3a 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -6,7 +6,7 @@ from typing import AsyncGenerator from alembic import command from alembic.config import Config from sqlalchemy.ext.asyncio.engine import create_async_engine -from sqlmodel import Field, Relationship, SQLModel, func, select +from sqlmodel import Field, Relationship, SQLModel, func, select, update from sqlmodel.ext.asyncio.session import AsyncSession from .logging import get_logger @@ -53,6 +53,14 @@ class ApiKey(SQLModel, table=True): # type: ignore return self.balance - self.reserved_balance +async def reset_all_reserved_balances(session: AsyncSession) -> None: + logger.info("Resetting all reserved balances to 0") + stmt = update(ApiKey).values(reserved_balance=0) + await session.exec(stmt) # type: ignore[call-overload] + await session.commit() + logger.info("Reserved balances reset successfully") + + class ModelRow(SQLModel, table=True): # type: ignore __tablename__ = "models" id: str = Field(primary_key=True) diff --git a/routstr/core/main.py b/routstr/core/main.py index 6ba783df..5bce5d32 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -61,6 +61,10 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: # Initialize application settings (env -> computed -> DB precedence) async with create_session() as session: s = await SettingsService.initialize(session) + if s.reset_reserved_balance_on_startup: + from .db import reset_all_reserved_balances + + await reset_all_reserved_balances(session) # Apply app metadata from settings try: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 52f6fae4..685f8970 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -61,6 +61,9 @@ class Settings(BaseSettings): tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE") # Minimum per-request charge in millisatoshis when model pricing is free/zero min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT") + reset_reserved_balance_on_startup: bool = Field( + default=True, env="RESET_RESERVED_BALANCE_ON_STARTUP" + ) # deactivate in horizontal scaling setups # Network cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") From 9229b87b7042dddc14ad7663e016c47d3b0e6f4a Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 10 Jan 2026 17:37:05 +0800 Subject: [PATCH 11/19] optimize price fetching --- routstr/payment/price.py | 26 ++++++++++++++++++++------ 1 file changed, 20 insertions(+), 6 deletions(-) diff --git a/routstr/payment/price.py b/routstr/payment/price.py index 9011ce1d..ad614322 100644 --- a/routstr/payment/price.py +++ b/routstr/payment/price.py @@ -79,15 +79,29 @@ async def _fetch_btc_usd_price() -> float: """Fetch the lowest BTC/USD price from multiple exchanges.""" async with httpx.AsyncClient(timeout=30.0) as client: try: - prices = await asyncio.gather( - _kraken_btc_usd(client), - _coinbase_btc_usd(client), - _binance_btc_usdt(client), - ) - valid_prices = [price for price in prices if price is not None] + tasks = [ + asyncio.create_task(_kraken_btc_usd(client)), + asyncio.create_task(_coinbase_btc_usd(client)), + asyncio.create_task(_binance_btc_usdt(client)), + ] + valid_prices: list[float] = [] + + for future in asyncio.as_completed(tasks): + price = await future + if price is not None: + valid_prices.append(price) + + if len(valid_prices) >= 2: + break + + for task in tasks: + if not task.done(): + task.cancel() + if not valid_prices: logger.error("No valid BTC prices obtained from any exchange") raise ValueError("Unable to fetch BTC price from any exchange") + return min(valid_prices) except Exception as e: logger.error( From 3bc38937e865b30d21557e814e52b972b611f83b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 18:39:24 +0100 Subject: [PATCH 12/19] ignore disabled provider --- routstr/payment/models.py | 7 +++++++ routstr/proxy.py | 2 ++ routstr/upstream/helpers.py | 2 ++ 3 files changed, 11 insertions(+) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 2a167838..4afbf45e 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -253,6 +253,11 @@ async def list_models( else 1.01, ) for r in rows + if include_disabled + or ( + r.upstream_provider_id in providers_by_id + and providers_by_id[r.upstream_provider_id].enabled + ) ] @@ -265,6 +270,8 @@ async def get_model_by_id( if not row or not row.enabled: return None provider = await session.get(UpstreamProviderRow, provider_id) + if not provider or not provider.enabled: + return None provider_fee = provider.provider_fee if provider else 1.01 return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee) diff --git a/routstr/proxy.py b/routstr/proxy.py index 1d5aaa96..1f9a48ba 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -98,6 +98,8 @@ async def refresh_model_maps() -> None: disabled_model_ids: set[str] = set() for provider in provider_rows: + if not provider.enabled: + continue for model in provider.models: if model.enabled: overrides_by_id[model.id] = (model, provider.provider_fee) diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index af333522..fdf267ab 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -102,6 +102,8 @@ async def get_all_models_with_overrides( ) for row in override_rows if row.upstream_provider_id is not None + and row.upstream_provider_id in providers_by_id + and providers_by_id[row.upstream_provider_id].enabled } all_models: dict[str, Model] = {} From fc042c768c9eec3a2e64c1f30af52594eedb4380 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 19:01:52 +0100 Subject: [PATCH 13/19] support child keys mapped to parent balance --- migrations/versions/a86e5348850b_.py | 42 ++++++ routstr/auth.py | 206 +++++++++++++++++++++------ routstr/balance.py | 94 ++++++++++-- routstr/core/admin.py | 3 +- routstr/core/db.py | 3 + routstr/core/settings.py | 1 + 6 files changed, 296 insertions(+), 53 deletions(-) create mode 100644 migrations/versions/a86e5348850b_.py diff --git a/migrations/versions/a86e5348850b_.py b/migrations/versions/a86e5348850b_.py new file mode 100644 index 00000000..12c35e41 --- /dev/null +++ b/migrations/versions/a86e5348850b_.py @@ -0,0 +1,42 @@ +""" + +Revision ID: a86e5348850b +Revises: b9667ffc5701 +Create Date: 2026-01-10 18:57:48.475781 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "a86e5348850b" +down_revision = "b9667ffc5701" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # Use batch_alter_table for SQLite compatibility + with op.batch_alter_table("api_keys", schema=None) as batch_op: + batch_op.add_column( + sa.Column( + "parent_key_hash", sqlmodel.sql.sqltypes.AutoString(), nullable=True + ) + ) + batch_op.create_index( + batch_op.f("ix_api_keys_parent_key_hash"), ["parent_key_hash"], unique=False + ) + batch_op.create_foreign_key( + "fk_api_keys_parent_key_hash", + "api_keys", + ["parent_key_hash"], + ["hashed_key"], + ) + + +def downgrade() -> None: + with op.batch_alter_table("api_keys", schema=None) as batch_op: + batch_op.drop_constraint("fk_api_keys_parent_key_hash", type_="foreignkey") + batch_op.drop_index(batch_op.f("ix_api_keys_parent_key_hash")) + batch_op.drop_column("parent_key_hash") diff --git a/routstr/auth.py b/routstr/auth.py index b3be04b2..1f869886 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -286,30 +286,55 @@ async def validate_bearer_key( ) +async def get_billing_key(key: ApiKey, session: AsyncSession) -> ApiKey: + """Returns the key that should be charged for the request.""" + if key.parent_key_hash: + parent = await session.get(ApiKey, key.parent_key_hash) + if parent: + # We want to keep the total_requests and total_spent on the child key + # but use the balance and reserved_balance of the parent. + # However, pay_for_request updates reserved_balance and total_requests. + # To stay simple, we charge the parent's balance and update parent's total_requests. + return parent + else: + logger.error( + "Parent key not found for child key", + extra={ + "child_key_hash": key.hashed_key[:8] + "...", + "parent_key_hash": key.parent_key_hash[:8] + "...", + }, + ) + return key + + async def pay_for_request( key: ApiKey, cost_per_request: int, session: AsyncSession ) -> int: """Process payment for a request.""" + billing_key = await get_billing_key(key, session) + logger.info( "Processing payment for request", extra={ "key_hash": key.hashed_key[:8] + "...", - "current_balance": key.balance, + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "current_balance": billing_key.balance, "required_cost": cost_per_request, - "sufficient_balance": key.balance >= cost_per_request, + "sufficient_balance": billing_key.balance >= cost_per_request, }, ) - if key.total_balance < cost_per_request: + if billing_key.total_balance < cost_per_request: logger.warning( "Insufficient balance for request", extra={ "key_hash": key.hashed_key[:8] + "...", - "balance": key.balance, - "reserved_balance": key.reserved_balance, + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "balance": billing_key.balance, + "reserved_balance": billing_key.reserved_balance, "required": cost_per_request, - "shortfall": cost_per_request - key.total_balance, + "shortfall": cost_per_request - billing_key.total_balance, }, ) @@ -317,7 +342,7 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.total_balance} available. (reserved: {key.reserved_balance})", + "message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.total_balance} available. (reserved: {billing_key.reserved_balance})", "type": "insufficient_quota", "code": "insufficient_balance", } @@ -328,15 +353,16 @@ async def pay_for_request( "Charging base cost for request", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "cost": cost_per_request, - "balance_before": key.balance, + "balance_before": billing_key.balance, }, ) # Charge the base cost for the request atomically to avoid race conditions stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= cost_per_request) .values( reserved_balance=col(ApiKey.reserved_balance) + cost_per_request, @@ -344,6 +370,16 @@ async def pay_for_request( ) ) result = await session.exec(stmt) # type: ignore[call-overload] + + # Also increment total_requests on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_requests=col(ApiKey.total_requests) + 1) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: @@ -351,8 +387,9 @@ async def pay_for_request( "Concurrent request depleted balance", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "required_cost": cost_per_request, - "current_balance": key.balance, + "current_balance": billing_key.balance, }, ) @@ -361,23 +398,26 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {key.balance} available.", + "message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.balance} available.", "type": "insufficient_quota", "code": "insufficient_balance", } }, ) - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Payment processed successfully", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": cost_per_request, - "new_balance": key.balance, - "total_spent": key.total_spent, - "total_requests": key.total_requests, + "new_balance": billing_key.balance, + "total_spent": billing_key.total_spent, + "total_requests": billing_key.total_requests, }, ) @@ -387,9 +427,11 @@ async def pay_for_request( async def revert_pay_for_request( key: ApiKey, session: AsyncSession, cost_per_request: int ) -> None: + billing_key = await get_billing_key(key, session) + stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, total_requests=col(ApiKey.total_requests) - 1, @@ -397,27 +439,40 @@ async def revert_pay_for_request( ) result = await session.exec(stmt) # type: ignore[call-overload] + + # Also decrement total_requests on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_requests=col(ApiKey.total_requests) - 1) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: logger.error( "Failed to revert payment - insufficient reserved balance", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "cost_to_revert": cost_per_request, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, }, ) raise HTTPException( status_code=402, detail={ "error": { - "message": f"failed to revert request payment: {cost_per_request} mSats required. {key.balance} available.", + "message": f"failed to revert request payment: {cost_per_request} mSats required. {billing_key.balance} available.", "type": "payment_error", "code": "payment_error", } }, ) - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) async def adjust_payment_for_tokens( @@ -428,15 +483,17 @@ async def adjust_payment_for_tokens( This is called after the initial payment and the upstream request is complete. Returns cost data to be included in the response. """ + billing_key = await get_billing_key(key, session) model = response_data.get("model", "unknown") logger.debug( "Starting payment adjustment for tokens", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "deducted_max_cost": deducted_max_cost, - "current_balance": key.balance, + "current_balance": billing_key.balance, "has_usage": "usage" in response_data, }, ) @@ -446,8 +503,10 @@ async def adjust_payment_for_tokens( try: release_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values(reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .values( + reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost + ) ) await session.exec(release_stmt) # type: ignore[call-overload] await session.commit() @@ -455,13 +514,18 @@ async def adjust_payment_for_tokens( "Released reservation without charging (fallback)", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, }, ) except Exception as e: logger.error( "Failed to release reservation in fallback", - extra={"error": str(e), "key_hash": key.hashed_key[:8] + "..."}, + extra={ + "error": str(e), + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + }, ) match await calculate_cost(response_data, deducted_max_cost, session): @@ -470,6 +534,7 @@ async def adjust_payment_for_tokens( "Using max cost data (no token adjustment)", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "max_cost": cost.total_msats, }, @@ -477,7 +542,7 @@ async def adjust_payment_for_tokens( # Finalize by releasing reservation and charging max cost finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, balance=col(ApiKey.balance) - cost.total_msats, @@ -485,27 +550,41 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + cost.total_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: logger.error( "Failed to finalize max-cost payment - retrying reservation release", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, "total_cost": cost.total_msats, "model": model, }, ) await release_reservation_only() else: - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Max cost payment finalized", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": cost.total_msats, - "new_balance": key.balance, + "new_balance": billing_key.balance, "model": model, }, ) @@ -521,6 +600,7 @@ async def adjust_payment_for_tokens( "Calculated token-based cost", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "model": model, "token_cost": cost.total_msats, "deducted_max_cost": deducted_max_cost, @@ -533,11 +613,15 @@ async def adjust_payment_for_tokens( if cost_difference == 0: logger.debug( "Finalizing with exact reserved cost", - extra={"key_hash": key.hashed_key[:8] + "...", "model": model}, + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "model": model, + }, ) finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -546,8 +630,20 @@ async def adjust_payment_for_tokens( ) ) await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) return cost.dict() # this should never happen why do we handle this??? @@ -557,16 +653,17 @@ async def adjust_payment_for_tokens( "Additional charge required for token usage", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "additional_charge": cost_difference, - "current_balance": key.balance, - "sufficient_balance": key.balance >= cost_difference, + "current_balance": billing_key.balance, + "sufficient_balance": billing_key.balance >= cost_difference, "model": model, }, ) finalize_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -575,18 +672,31 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(finalize_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount: cost.total_msats = total_cost_msats - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Finalized payment with additional charge", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "charged_amount": total_cost_msats, - "new_balance": key.balance, + "new_balance": billing_key.balance, "model": model, }, ) @@ -595,6 +705,7 @@ async def adjust_payment_for_tokens( "Failed to finalize additional charge - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "attempted_charge": total_cost_msats, "model": model, }, @@ -607,15 +718,16 @@ async def adjust_payment_for_tokens( "Refunding excess payment", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "refund_amount": refund, - "current_balance": key.balance, + "current_balance": billing_key.balance, "model": model, }, ) refund_stmt = ( update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.hashed_key) == billing_key.hashed_key) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, @@ -624,6 +736,16 @@ async def adjust_payment_for_tokens( ) ) result = await session.exec(refund_stmt) # type: ignore[call-overload] + + # Also update total_spent on the child key if it's different + if billing_key.hashed_key != key.hashed_key: + child_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values(total_spent=col(ApiKey.total_spent) + total_cost_msats) + ) + await session.exec(child_stmt) # type: ignore[call-overload] + await session.commit() if result.rowcount == 0: @@ -631,8 +753,9 @@ async def adjust_payment_for_tokens( "Failed to finalize payment - releasing reservation", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": key.reserved_balance, + "current_reserved_balance": billing_key.reserved_balance, "total_cost": total_cost_msats, "model": model, }, @@ -640,14 +763,17 @@ async def adjust_payment_for_tokens( await release_reservation_only() else: cost.total_msats = total_cost_msats - await session.refresh(key) + await session.refresh(billing_key) + if billing_key.hashed_key != key.hashed_key: + await session.refresh(key) logger.info( "Refund processed successfully", extra={ "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", "refunded_amount": refund, - "new_balance": key.balance, + "new_balance": billing_key.balance, "final_cost": cost.total_msats, "model": model, }, diff --git a/routstr/balance.py b/routstr/balance.py index 6697a5db..472f8b05 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -32,16 +32,30 @@ async def get_key_from_header( ) -# TODO: remove this endpoint when frontend is updated -@router.get("/", include_in_schema=False) -async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict: +async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict: + from .auth import get_billing_key + + billing_key = await get_billing_key(key, session) return { "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - "reserved": key.reserved_balance, + "balance": billing_key.balance, + "reserved": billing_key.reserved_balance, + "is_child": key.parent_key_hash is not None, + "parent_key": "sk-" + key.parent_key_hash if key.parent_key_hash else None, + "total_requests": key.total_requests, + "total_spent": key.total_spent, } +# TODO: remove this endpoint when frontend is updated +@router.get("/", include_in_schema=False) +async def account_info( + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + return await get_balance_info(key, session) + + # TODO: Implement POST /v1/wallet/create endpoint # This endpoint should accept: # - cashu_token (required): The eCash token to deposit @@ -66,12 +80,11 @@ async def create_balance( @router.get("/info") -async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict: - return { - "api_key": "sk-" + key.hashed_key, - "balance": key.balance, - "reserved": key.reserved_balance, - } +async def wallet_info( + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + return await get_balance_info(key, session) class TopupRequest(BaseModel): @@ -85,6 +98,10 @@ async def topup_wallet_endpoint( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict[str, int]: + from .auth import get_billing_key + + billing_key = await get_billing_key(key, session) + if topup_request is not None: cashu_token = topup_request.cashu_token if cashu_token is None: @@ -94,7 +111,7 @@ async def topup_wallet_endpoint( if len(cashu_token) < 10 or "cashu" not in cashu_token: raise HTTPException(status_code=400, detail="Invalid token format") try: - amount_msats = await credit_balance(cashu_token, key, session) + amount_msats = await credit_balance(cashu_token, billing_key, session) except ValueError as e: error_msg = str(e) if "already spent" in error_msg.lower(): @@ -155,6 +172,12 @@ async def refund_wallet_endpoint( key: ApiKey = await validate_bearer_key(bearer_value, session) + if key.parent_key_hash: + raise HTTPException( + status_code=400, + detail="Cannot refund child key. Please refund the parent key instead.", + ) + remaining_balance_msats: int = key.total_balance if key.refund_currency == "sat": @@ -240,6 +263,53 @@ async def wallet_catch_all(path: str) -> NoReturn: ) +@router.post("/child-key") +async def create_child_key( + key: ApiKey = Depends(get_key_from_header), + session: AsyncSession = Depends(get_session), +) -> dict: + """Creates a child API key that uses the parent's balance.""" + # Check if this is already a child key + if key.parent_key_hash: + raise HTTPException( + status_code=400, + detail="Cannot create a child key for another child key.", + ) + + cost = settings.child_key_cost + + if key.total_balance < cost: + raise HTTPException( + status_code=402, + detail=f"Insufficient balance to create child key. {cost} mSats required.", + ) + + # Deduct cost from parent + key.balance -= cost + key.total_spent += cost + session.add(key) + + # Generate new key + import secrets + + new_key_raw = secrets.token_hex(32) + new_key_hash = new_key_raw # We use the raw key as the hash for sk- keys + + child_key = ApiKey( + hashed_key=new_key_hash, + balance=0, + parent_key_hash=key.hashed_key, + ) + session.add(child_key) + await session.commit() + + return { + "api_key": "sk-" + new_key_hash, + "cost_msats": cost, + "parent_balance": key.balance, + } + + balance_router.include_router(lightning_router) balance_router.include_router(router) diff --git a/routstr/core/admin.py b/routstr/core/admin.py index d4f7ec59..626df30c 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -124,7 +124,7 @@ async def partial_apikeys(request: Request) -> str: rows = "".join( [ - f"{key.hashed_key}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{fmt_time(key.key_expiry_time)}" + f"{key.hashed_key}{'
(Child of ' + key.parent_key_hash[:8] + '...)' if key.parent_key_hash else ''}{key.balance}{key.total_spent}{key.total_requests}{key.refund_address}{fmt_time(key.key_expiry_time)}" for key in api_keys ] ) @@ -158,6 +158,7 @@ async def get_temporary_balances_api(request: Request) -> list[dict[str, object] "total_requests": key.total_requests, "refund_address": key.refund_address, "key_expiry_time": key.key_expiry_time, + "parent_key_hash": key.parent_key_hash, } for key in api_keys ] diff --git a/routstr/core/db.py b/routstr/core/db.py index 4c236d3a..46c56582 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -47,6 +47,9 @@ class ApiKey(SQLModel, table=True): # type: ignore default=None, description="Currency of the cashu-token", ) + parent_key_hash: str | None = Field( + default=None, foreign_key="api_keys.hashed_key", index=True + ) @property def total_balance(self) -> int: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 685f8970..3b659615 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -59,6 +59,7 @@ class Settings(BaseSettings): exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE") upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE") tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE") + child_key_cost: int = Field(default=1000, env="CHILD_KEY_COST") # Minimum per-request charge in millisatoshis when model pricing is free/zero min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT") reset_reserved_balance_on_startup: bool = Field( From daf17f51ab79d60fda29dfb33f18a67a8571e357 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 20:00:52 +0100 Subject: [PATCH 14/19] fix: move child-key route before catch-all and fix indentation --- examples/create_child_keys.py | 44 +++++++++ routstr/balance.py | 24 ++--- routstr/core/main.py | 8 +- tests/integration/test_child_keys.py | 131 +++++++++++++++++++++++++++ 4 files changed, 190 insertions(+), 17 deletions(-) create mode 100644 examples/create_child_keys.py create mode 100644 tests/integration/test_child_keys.py diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py new file mode 100644 index 00000000..eefa3710 --- /dev/null +++ b/examples/create_child_keys.py @@ -0,0 +1,44 @@ +import httpx +import sys +import json + + +def create_child_keys(base_url, api_key, count=3): + headers = {"Authorization": f"Bearer {api_key}"} + + print(f"Requesting {count} child keys from {base_url}...") + + child_keys = [] + + for i in range(count): + try: + response = httpx.post(f"{base_url}/v1/balance/child-key", headers=headers) + if response.status_code == 200: + data = response.json() + child_keys.append(data["api_key"]) + print( + f" [{i + 1}] Created: {data['api_key']} (Cost: {data['cost_msats']} msats)" + ) + else: + print(f" [{i + 1}] Failed: {response.status_code} - {response.text}") + except Exception as e: + print(f" [{i + 1}] Error: {str(e)}") + + return child_keys + + +if __name__ == "__main__": + if len(sys.argv) < 2: + print("Usage: python create_child_keys.py [base_url]") + sys.exit(1) + + auth_key = sys.argv[1] + base_url = sys.argv[2] if len(sys.argv) > 2 else "http://localhost:8000" + + keys = create_child_keys(base_url, auth_key) + + if keys: + print("\nSuccessfully created child keys:") + print(json.dumps(keys, indent=2)) + else: + print("\nNo child keys were created.") diff --git a/routstr/balance.py b/routstr/balance.py index 472f8b05..3ee8ca11 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -251,18 +251,6 @@ async def donate(token: str, ref: str | None = None) -> str: return "Invalid token." -@router.api_route( - "/{path:path}", - methods=["GET", "POST", "PUT", "DELETE"], - include_in_schema=False, - response_model=None, -) -async def wallet_catch_all(path: str) -> NoReturn: - raise HTTPException( - status_code=404, detail="Not found check /docs for available endpoints" - ) - - @router.post("/child-key") async def create_child_key( key: ApiKey = Depends(get_key_from_header), @@ -310,6 +298,18 @@ async def create_child_key( } +@router.api_route( + "/{path:path}", + methods=["GET", "POST", "PUT", "DELETE"], + include_in_schema=False, + response_model=None, +) +async def wallet_catch_all(path: str) -> NoReturn: + raise HTTPException( + status_code=404, detail="Not found check /docs for available endpoints" + ) + + balance_router.include_router(lightning_router) balance_router.include_router(router) diff --git a/routstr/core/main.py b/routstr/core/main.py index 5bce5d32..3083097d 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -13,12 +13,10 @@ from starlette.exceptions import HTTPException from ..balance import balance_router, deprecated_wallet_router from ..discovery import providers_cache_refresher, providers_router from ..nip91 import announce_provider -from ..payment.models import ( - models_router, - update_sats_pricing, -) +from ..payment.models import models_router, update_sats_pricing from ..payment.price import update_prices_periodically -from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically +from ..proxy import (initialize_upstreams, proxy_router, + refresh_model_maps_periodically) from ..wallet import periodic_payout from .admin import admin_router from .db import create_session, init_db, run_migrations diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py new file mode 100644 index 00000000..77da7675 --- /dev/null +++ b/tests/integration/test_child_keys.py @@ -0,0 +1,131 @@ +import pytest +import secrets +from fastapi import HTTPException +from routstr.core.db import ApiKey +from routstr.balance import create_child_key +from routstr.auth import pay_for_request, adjust_payment_for_tokens +from routstr.core.settings import settings + + +@pytest.mark.asyncio +async def test_child_key_flow(integration_session): + # 1. Create a parent key with balance + parent_raw = "parent_test_key_" + secrets.token_hex(4) + parent_key = ApiKey( + hashed_key=parent_raw, + balance=10000, # 10 sats + ) + integration_session.add(parent_key) + await integration_session.commit() + await integration_session.refresh(parent_key) + + # Mock settings + settings.child_key_cost = 1000 # 1 sat + + # 2. Call create_child_key + result = await create_child_key(parent_key, integration_session) + + assert "api_key" in result + assert result["cost_msats"] == 1000 + assert result["parent_balance"] == 9000 + + child_key_raw = result["api_key"][3:] # remove sk- + + # 3. Verify child key exists in DB + child_key_db = await integration_session.get(ApiKey, child_key_raw) + assert child_key_db is not None + assert child_key_db.parent_key_hash == parent_key.hashed_key + assert child_key_db.balance == 0 + + # 4. Test payment with child key + cost = 500 + await pay_for_request(child_key_db, cost, integration_session) + + # Refresh keys + await integration_session.refresh(parent_key) + await integration_session.refresh(child_key_db) + + # Parent should be charged + assert parent_key.reserved_balance == 500 + assert parent_key.total_requests == 1 + + # Child should have total_requests incremented + assert child_key_db.total_requests == 1 + + # 5. Test adjustment + response_data = {"model": "test-model", "usage": {"total_tokens": 10}} + + # Mock calculate_cost + import routstr.auth + from routstr.payment.cost_calculation import CostData + + async def mock_calculate_cost(*args, **kwargs): + return CostData( + base_msats=0, input_msats=200, output_msats=200, total_msats=400 + ) + + # Patch calculate_cost + original_calculate_cost = routstr.auth.calculate_cost + routstr.auth.calculate_cost = mock_calculate_cost + + try: + adjustment = await adjust_payment_for_tokens( + child_key_db, response_data, integration_session, 500 + ) + assert adjustment["total_msats"] == 400 + + # Refresh keys + await integration_session.refresh(parent_key) + await integration_session.refresh(child_key_db) + + # Parent should have updated balance and total_spent + assert parent_key.reserved_balance == 0 + assert parent_key.balance == 9000 - 400 + assert ( + parent_key.total_spent == 1400 + ) # 1000 for child key creation + 400 for request + + # Child should also have total_spent updated + assert child_key_db.total_spent == 400 + + finally: + routstr.auth.calculate_cost = original_calculate_cost + + +@pytest.mark.asyncio +async def test_child_key_insufficient_balance(integration_session): + parent_key = ApiKey( + hashed_key="poor_parent_" + secrets.token_hex(4), + balance=500, + ) + integration_session.add(parent_key) + await integration_session.commit() + await integration_session.refresh(parent_key) + + settings.child_key_cost = 1000 + + with pytest.raises(HTTPException) as exc: + await create_child_key(parent_key, integration_session) + assert exc.value.status_code == 402 + + +@pytest.mark.asyncio +async def test_child_key_cannot_create_child(integration_session): + parent_key = ApiKey( + hashed_key="parent_" + secrets.token_hex(4), + balance=10000, + ) + child_key = ApiKey( + hashed_key="child_" + secrets.token_hex(4), + balance=0, + parent_key_hash=parent_key.hashed_key, + ) + integration_session.add(parent_key) + integration_session.add(child_key) + await integration_session.commit() + await integration_session.refresh(child_key) + + with pytest.raises(HTTPException) as exc: + await create_child_key(child_key, integration_session) + assert exc.value.status_code == 400 + assert "Cannot create a child key for another child key" in str(exc.value.detail) From 917a4d32b1630440a6f4ed00fb547ba55a908c07 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 20:10:59 +0100 Subject: [PATCH 15/19] ui: visualize parent-child relationship in balances page --- ui/components/temporary-balances.tsx | 79 +++++++++++++++++++++++----- ui/lib/api/services/admin.ts | 1 + 2 files changed, 67 insertions(+), 13 deletions(-) diff --git a/ui/components/temporary-balances.tsx b/ui/components/temporary-balances.tsx index c33d0d43..004ba77f 100644 --- a/ui/components/temporary-balances.tsx +++ b/ui/components/temporary-balances.tsx @@ -60,7 +60,11 @@ export function TemporaryBalances({ let totalRequests = 0; balances.forEach((balance) => { - totalBalance += balance.balance || 0; + // Only count parents for total balance to avoid double counting + // since child keys use parent balance + if (!balance.parent_key_hash) { + totalBalance += balance.balance || 0; + } totalSpent += balance.total_spent || 0; totalRequests += balance.total_requests || 0; }); @@ -72,6 +76,34 @@ export function TemporaryBalances({ ? calculateTotals(data) : { totalBalance: 0, totalSpent: 0, totalRequests: 0 }; + // Group parents and children + const hierarchicalData = (() => { + if (!data) return []; + + const parents = filteredData.filter((item) => !item.parent_key_hash); + const result: (TemporaryBalance & { isChild?: boolean })[] = []; + + parents.forEach((parent) => { + result.push(parent); + const children = data.filter( + (item) => item.parent_key_hash === parent.hashed_key + ); + children.forEach((child) => { + result.push({ ...child, isChild: true }); + }); + }); + + // Add children whose parents didn't match the search or aren't in the list + const orphans = filteredData.filter( + (item) => + item.parent_key_hash && + !result.some((r) => r.hashed_key === item.hashed_key) + ); + result.push(...orphans.map((o) => ({ ...o, isChild: true }))); + + return result; + })(); + return ( <> @@ -182,22 +214,32 @@ export function TemporaryBalances({
Expiry Time
- {filteredData.length > 0 ? ( - filteredData.map((balance, index) => ( + {hierarchicalData.length > 0 ? ( + hierarchicalData.map((balance, index) => (
{/* Desktop Layout */}
-
+
+ {balance.isChild && ( + + Child + + )} {balance.hashed_key}
- {formatBalance(balance.balance)} + {balance.isChild ? ( + (Parent) + ) : ( + formatBalance(balance.balance) + )}
{formatBalance(balance.total_spent)} @@ -226,13 +268,20 @@ export function TemporaryBalances({ {/* Mobile Layout */}
-
- - Key - -
- {balance.hashed_key} +
+
+ + {balance.isChild ? 'Child Key' : 'Key'} + +
+ {balance.hashed_key} +
+ {balance.isChild && ( + + Child + + )}
@@ -241,7 +290,11 @@ export function TemporaryBalances({ Balance
- {formatBalance(balance.balance)} + {balance.isChild ? ( + (Uses Parent) + ) : ( + formatBalance(balance.balance) + )}
diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index f94729ba..2ea27e7e 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -926,6 +926,7 @@ export const TemporaryBalanceSchema = z.object({ total_requests: z.number(), refund_address: z.string().nullable(), key_expiry_time: z.number().nullable(), + parent_key_hash: z.string().nullable().optional(), }); export type TemporaryBalance = z.infer; From 88bcc0edcb9cebc3532861ddfe816f33718c0d1c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 21:34:48 +0100 Subject: [PATCH 16/19] fmt --- examples/create_child_keys.py | 6 ++++-- routstr/core/main.py | 3 +-- tests/integration/test_child_keys.py | 8 +++++--- 3 files changed, 10 insertions(+), 7 deletions(-) diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py index eefa3710..5e5d23e7 100644 --- a/examples/create_child_keys.py +++ b/examples/create_child_keys.py @@ -1,6 +1,7 @@ -import httpx -import sys import json +import sys + +import httpx def create_child_keys(base_url, api_key, count=3): @@ -42,3 +43,4 @@ if __name__ == "__main__": print(json.dumps(keys, indent=2)) else: print("\nNo child keys were created.") + diff --git a/routstr/core/main.py b/routstr/core/main.py index 3083097d..00a51ec3 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -15,8 +15,7 @@ from ..discovery import providers_cache_refresher, providers_router from ..nip91 import announce_provider from ..payment.models import models_router, update_sats_pricing from ..payment.price import update_prices_periodically -from ..proxy import (initialize_upstreams, proxy_router, - refresh_model_maps_periodically) +from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..wallet import periodic_payout from .admin import admin_router from .db import create_session, init_db, run_migrations diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py index 77da7675..4265430a 100644 --- a/tests/integration/test_child_keys.py +++ b/tests/integration/test_child_keys.py @@ -1,9 +1,11 @@ -import pytest import secrets + +import pytest from fastapi import HTTPException -from routstr.core.db import ApiKey + +from routstr.auth import adjust_payment_for_tokens, pay_for_request from routstr.balance import create_child_key -from routstr.auth import pay_for_request, adjust_payment_for_tokens +from routstr.core.db import ApiKey from routstr.core.settings import settings From d4339287beb8173bfdf80f5202083fad297eb585 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 21:35:16 +0100 Subject: [PATCH 17/19] chore: add type annotations to example and test files --- examples/create_child_keys.py | 3 +-- tests/integration/test_child_keys.py | 12 ++++++++---- 2 files changed, 9 insertions(+), 6 deletions(-) diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py index 5e5d23e7..24f4556d 100644 --- a/examples/create_child_keys.py +++ b/examples/create_child_keys.py @@ -4,7 +4,7 @@ import sys import httpx -def create_child_keys(base_url, api_key, count=3): +def create_child_keys(base_url: str, api_key: str, count: int = 3) -> list[str]: headers = {"Authorization": f"Bearer {api_key}"} print(f"Requesting {count} child keys from {base_url}...") @@ -43,4 +43,3 @@ if __name__ == "__main__": print(json.dumps(keys, indent=2)) else: print("\nNo child keys were created.") - diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py index 4265430a..e868d1cb 100644 --- a/tests/integration/test_child_keys.py +++ b/tests/integration/test_child_keys.py @@ -1,7 +1,9 @@ import secrets +from typing import Any import pytest from fastapi import HTTPException +from sqlmodel.ext.asyncio.session import AsyncSession from routstr.auth import adjust_payment_for_tokens, pay_for_request from routstr.balance import create_child_key @@ -10,7 +12,7 @@ from routstr.core.settings import settings @pytest.mark.asyncio -async def test_child_key_flow(integration_session): +async def test_child_key_flow(integration_session: AsyncSession) -> None: # 1. Create a parent key with balance parent_raw = "parent_test_key_" + secrets.token_hex(4) parent_key = ApiKey( @@ -61,7 +63,7 @@ async def test_child_key_flow(integration_session): import routstr.auth from routstr.payment.cost_calculation import CostData - async def mock_calculate_cost(*args, **kwargs): + async def mock_calculate_cost(*args: Any, **kwargs: Any) -> CostData: return CostData( base_msats=0, input_msats=200, output_msats=200, total_msats=400 ) @@ -95,7 +97,9 @@ async def test_child_key_flow(integration_session): @pytest.mark.asyncio -async def test_child_key_insufficient_balance(integration_session): +async def test_child_key_insufficient_balance( + integration_session: AsyncSession, +) -> None: parent_key = ApiKey( hashed_key="poor_parent_" + secrets.token_hex(4), balance=500, @@ -112,7 +116,7 @@ async def test_child_key_insufficient_balance(integration_session): @pytest.mark.asyncio -async def test_child_key_cannot_create_child(integration_session): +async def test_child_key_cannot_create_child(integration_session: AsyncSession) -> None: parent_key = ApiKey( hashed_key="parent_" + secrets.token_hex(4), balance=10000, From db021866d87e4f7e464195b85d4d62d88696146a Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 21:42:21 +0100 Subject: [PATCH 18/19] fmt u --- ui/components/temporary-balances.tsx | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/ui/components/temporary-balances.tsx b/ui/components/temporary-balances.tsx index 004ba77f..7e7927d9 100644 --- a/ui/components/temporary-balances.tsx +++ b/ui/components/temporary-balances.tsx @@ -220,15 +220,18 @@ export function TemporaryBalances({ key={index} className={cn( 'hover:bg-muted/50 border-t p-3 text-sm transition-colors', - balance.balance === 0 && !balance.isChild && 'opacity-60', - balance.isChild && 'bg-blue-50/30 ml-4 border-l-2 border-l-blue-200' + balance.balance === 0 && + !balance.isChild && + 'opacity-60', + balance.isChild && + 'ml-4 border-l-2 border-l-blue-200 bg-blue-50/30' )} > {/* Desktop Layout */}
{balance.isChild && ( - + Child )} @@ -236,7 +239,9 @@ export function TemporaryBalances({
{balance.isChild ? ( - (Parent) + + (Parent) + ) : ( formatBalance(balance.balance) )} @@ -278,7 +283,7 @@ export function TemporaryBalances({
{balance.isChild && ( - + Child )} @@ -291,7 +296,9 @@ export function TemporaryBalances({
{balance.isChild ? ( - (Uses Parent) + + (Uses Parent) + ) : ( formatBalance(balance.balance) )} From 24015ebec1cb4754c6f0b63b12dc008a272b93fa Mon Sep 17 00:00:00 2001 From: Shroominic Date: Tue, 13 Jan 2026 06:31:24 +0800 Subject: [PATCH 19/19] Revert "Merge remote-tracking branch 'origin/mock-upstream-with-testnut-mint' into v0.2.2" This reverts commit 5a4ba60072ede270eda6f3def55cf78ae5063f8b, reversing changes made to eed5bc5b0436660729fbeeb7a8cfc49783ca24cd. --- routstr/core/settings.py | 9 +- routstr/proxy.py | 20 --- routstr/upstream/fake.py | 265 ------------------------------------ routstr/upstream/helpers.py | 8 -- routstr/wallet.py | 2 - 5 files changed, 1 insertion(+), 303 deletions(-) delete mode 100644 routstr/upstream/fake.py diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 685f8970..7507a48e 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -6,7 +6,7 @@ import os from datetime import datetime, timezone from typing import Any -from pydantic.v1 import BaseModel, BaseSettings, Field, validator +from pydantic.v1 import BaseModel, BaseSettings, Field from sqlmodel.ext.asyncio.session import AsyncSession @@ -37,13 +37,6 @@ class Settings(BaseSettings): # Cashu cashu_mints: list[str] = Field(default_factory=list, env="CASHU_MINTS") - - @validator("cashu_mints", pre=True, each_item=True) - def normalize_mint_url(cls, v: str) -> str: - if isinstance(v, str): - return v.rstrip("/") - return v - receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS") primary_mint: str = Field(default="", env="PRIMARY_MINT_URL") primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT") diff --git a/routstr/proxy.py b/routstr/proxy.py index 1d5aaa96..58ef12d4 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -16,7 +16,6 @@ from .core.db import ( create_session, get_session, ) -from .core.settings import settings from .payment.helpers import ( calculate_discounted_max_cost, check_token_balance, @@ -26,7 +25,6 @@ from .payment.helpers import ( from .payment.models import Model from .upstream import BaseUpstreamProvider from .upstream.helpers import init_upstreams -from .wallet import deserialize_token_from_string logger = get_logger(__name__) proxy_router = APIRouter() @@ -148,24 +146,6 @@ async def proxy( else: model_id = request_body_dict.get("model", "unknown") - if "https://testnut.cashu.space" in settings.cashu_mints: - try: - token_str = None - if x_cashu_header := headers.get("x-cashu"): - token_str = x_cashu_header - elif auth_header := headers.get("authorization"): - parts = auth_header.split(" ") - if len(parts) > 1 and not parts[1].startswith("sk-"): - token_str = parts[1] - - if token_str: - token_obj = deserialize_token_from_string(token_str) - if token_obj.mint == "https://testnut.cashu.space": - model_id = "mock/gpt-420-mock" - request_body_dict["model"] = model_id - except Exception: - pass - model_obj = get_model_instance(model_id) if not model_obj: return create_error_response( diff --git a/routstr/upstream/fake.py b/routstr/upstream/fake.py deleted file mode 100644 index cfb530d1..00000000 --- a/routstr/upstream/fake.py +++ /dev/null @@ -1,265 +0,0 @@ -import asyncio -import json -import random -from typing import AsyncIterator - -from fastapi import Request -from fastapi.responses import Response, StreamingResponse - -from ..core.db import ApiKey, AsyncSession -from ..payment.models import Architecture, Model, Pricing -from .base import BaseUpstreamProvider - - -class MockUpstreamProvider(BaseUpstreamProvider): - """Fack Mock Upstream provider specifically for Testing.""" - - provider_type = "mock" - - async def forward_request( - self, - request: Request, - path: str, - headers: dict, - request_body: bytes | None, - key: ApiKey, - max_cost_for_model: int, - session: AsyncSession, - model_obj: Model, - ) -> Response | StreamingResponse: - if path.endswith("chat/completions"): - is_streaming = False - if request_body: - request_data = json.loads(request_body) - is_streaming = request_data.get("stream", False) - - if is_streaming: - - async def fake_streaming_response( - chunk_size: int | None = None, - ) -> AsyncIterator[bytes]: - suffix = random.randint(1000, 9999) - req_id = f"gen-mock-stream-{suffix}" - created = 1766138895 - model = "mock/gpt-420-mock" - - def make_chunk( - delta: dict, - finish_reason: str | None = None, - usage: dict | None = None, - ) -> bytes: - chunk = { - "id": req_id, - "provider": "MockProvider", - "model": model, - "object": "chat.completion.chunk", - "created": created, - "choices": [ - { - "index": 0, - "delta": delta, - "finish_reason": finish_reason, - "native_finish_reason": "completed" - if finish_reason - else None, - "logprobs": None, - } - ], - } - if usage: - chunk["usage"] = usage - return f"data: {json.dumps(chunk)}\n\n".encode() - - # 1. Initial chunk - yield make_chunk({"role": "assistant", "content": ""}) - await asyncio.sleep(0.02) - - # 2. Reasoning chunks - reasoning_tokens = ["Mock", " reason", "ing", "..."] - for token in reasoning_tokens: - delta = { - "role": "assistant", - "content": "", - "reasoning": token, - "reasoning_details": [ - { - "type": "reasoning.summary", - "summary": token, - "format": "openai-responses-v1", - "index": 0, - } - ], - } - yield make_chunk(delta) - await asyncio.sleep(0.03) - - # 3. Content chunks - content_tokens = ["This", " is", " a", " mock", " stream", "."] - for token in content_tokens: - yield make_chunk({"role": "assistant", "content": token}) - await asyncio.sleep(0.03) - - # 4. Finish chunk - yield make_chunk( - {"role": "assistant", "content": ""}, finish_reason="stop" - ) - - # 5. Usage chunk - usage_data = { - "prompt_tokens": 10, - "completion_tokens": 20, - "total_tokens": 30, - "cost": 0.001, - "is_byok": False, - "prompt_tokens_details": { - "cached_tokens": 0, - "audio_tokens": 0, - "video_tokens": 0, - }, - "cost_details": { - "upstream_inference_cost": None, - "upstream_inference_prompt_cost": 0, - "upstream_inference_completions_cost": 0.001, - }, - "completion_tokens_details": { - "reasoning_tokens": 10, - "image_tokens": 0, - }, - } - - usage_chunk = { - "id": req_id, - "provider": "MockProvider", - "model": model, - "object": "chat.completion.chunk", - "created": created, - "choices": [ - { - "index": 0, - "delta": {"role": "assistant", "content": ""}, - "finish_reason": None, - "native_finish_reason": None, - "logprobs": None, - } - ], - "usage": usage_data, - } - yield f"data: {json.dumps(usage_chunk)}\n\n".encode() - - # 6. DONE - yield b"data: [DONE]\n\n" - - # 7. Cost - cost_chunk = { - "cost": { - "base_msats": 0, - "input_msats": 2, - "output_msats": 10, - "total_msats": 12, - } - } - yield f"data: {json.dumps(cost_chunk)}\n\n".encode() - - return StreamingResponse( - fake_streaming_response(), - 200, - ) - - else: - suffix = random.randint(1000, 9999) - content_dict = { - "id": f"gen-mock-{suffix}", - "provider": "MockProvider", - "model": "mock/gpt-5-mini", - "object": "chat.completion", - "created": 1766138655, - "choices": [ - { - "logprobs": None, - "finish_reason": "length", - "native_finish_reason": "max_output_tokens", - "index": 0, - "message": { - "role": "assistant", - "content": f"Mock Content {suffix}", - "refusal": None, - "reasoning": f"Mock Reasoning {suffix}", - "reasoning_details": [ - { - "format": "openai-responses-v1", - "index": 0, - "type": "reasoning.summary", - "summary": f"Mock Summary {suffix}", - }, - { - "id": f"rs_mock_{suffix}", - "format": "openai-responses-v1", - "index": 0, - "type": "reasoning.encrypted", - "data": "mock_encrypted_data", - }, - ], - }, - } - ], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 10, - "total_tokens": 20, - "cost": 0, - "is_byok": False, - "prompt_tokens_details": { - "cached_tokens": 0, - "audio_tokens": 0, - "video_tokens": 0, - }, - "cost_details": { - "upstream_inference_cost": None, - "upstream_inference_prompt_cost": 0, - "upstream_inference_completions_cost": 0, - }, - "completion_tokens_details": { - "reasoning_tokens": 5, - "image_tokens": 0, - }, - }, - "cost": { - "base_msats": 0, - "input_msats": 0, - "output_msats": 0, - "total_msats": 0, - }, - } - return Response(json.dumps(content_dict).encode(), 200) - - elif path.endswith("embeddings"): - raise NotImplementedError - elif path.endswith("responses"): - raise NotImplementedError - else: - raise NotImplementedError - - async def fetch_models(self) -> list[Model]: - return [ - Model( - id="mock/gpt-420-mock", - name="mock/gpt-420-mock", - created=0, - description="mock model for testing", - context_length=8192, - architecture=Architecture( - modality="text", - input_modalities=["text"], - output_modalities=["text"], - tokenizer="", - instruct_type=None, - ), - pricing=Pricing(prompt=0.01, completion=0.01), - ), - ] - - def transform_model_name(self, model_id: str) -> str: - return "fake-model" - - async def get_balance(self) -> float | None: - return 420.69 diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index af333522..c2417d40 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -218,14 +218,6 @@ async def init_upstreams() -> list[BaseUpstreamProvider]: results = await asyncio.gather(*tasks) upstreams = [p for p in results if p is not None] - if "https://testnut.cashu.space" in settings.cashu_mints: - from .fake import MockUpstreamProvider - - mock_provider = MockUpstreamProvider("mock", "mock") - await mock_provider.refresh_models_cache() - upstreams.append(mock_provider) - logger.info("Initialized MockUpstreamProvider for testnut mint") - return upstreams diff --git a/routstr/wallet.py b/routstr/wallet.py index 7e8981b2..ce87b560 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -313,8 +313,6 @@ async def periodic_payout() -> None: try: async with db.create_session() as session: for mint_url in settings.cashu_mints: - if mint_url == "https://testnut.cashu.space": - continue for unit in ["sat", "msat"]: wallet = await get_wallet(mint_url, unit) proofs = get_proofs_per_mint_and_unit(