diff --git a/pyproject.toml b/pyproject.toml
index f192541e..de7b7d2e 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -22,6 +22,9 @@ dev = [
"pytest-asyncio>=0.24.0",
"pytest-cov>=6.1.1",
"httpx>=0.25.2",
+ "psutil>=5.9.0",
+ "aiohttp>=3.9.0",
+ "pytest-benchmark>=4.0.0",
]
[tool.pytest.ini_options]
@@ -41,8 +44,10 @@ addopts = [
]
markers = [
"asyncio: marks tests as async (deselect with '-m \"not asyncio\"')",
- "integration: marks tests as integration tests",
+ "integration: marks tests as integration tests (deselect with '-m \"not integration\"')",
"unit: marks tests as unit tests",
+ "slow: marks tests as slow running (deselect with '-m \"not slow\"')",
+ "requires_real_mint: marks tests that require a running Cashu mint instance",
]
[tool.ruff.lint]
diff --git a/router/account.py b/router/account.py
index f9b4e6b9..1d1a0880 100644
--- a/router/account.py
+++ b/router/account.py
@@ -50,7 +50,63 @@ async def topup_wallet_endpoint(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict[str, int]:
- amount_msats = await credit_balance(cashu_token, key, session)
+ # Validate token format first
+ if not cashu_token or not cashu_token.startswith("cashu"):
+ raise HTTPException(status_code=400, detail="Invalid token format")
+
+ # Check for obviously invalid tokens
+ if len(cashu_token) < 10: # Too short to be valid
+ raise HTTPException(status_code=400, detail="Invalid token format")
+
+ # Check for malformed base64 in token
+ if "cashuA" in cashu_token:
+ try:
+ import base64
+
+ # Extract base64 part after 'cashuA'
+ base64_part = cashu_token[6:]
+ if base64_part:
+ # Try to decode - will raise exception if invalid
+ base64.urlsafe_b64decode(base64_part + "=" * (4 - len(base64_part) % 4))
+ except Exception:
+ raise HTTPException(
+ status_code=400, detail="Invalid token format: malformed base64"
+ )
+
+ # Check for newlines or other invalid characters
+ if any(char in cashu_token for char in ["\n", "\r", "\t"]):
+ raise HTTPException(status_code=400, detail="Invalid token format")
+
+ # Capture stdout to detect errors from credit_balance
+ import io
+ from contextlib import redirect_stdout
+
+ f = io.StringIO()
+ with redirect_stdout(f):
+ amount_msats = await credit_balance(cashu_token, key, session)
+
+ output = f.getvalue()
+
+ # Check for errors in the output
+ if "Error in credit_balance:" in output and amount_msats == 0:
+ error_msg = output.split("Error in credit_balance: ")[-1].strip()
+ # Common error patterns
+ if "Token already spent" in error_msg:
+ raise HTTPException(status_code=400, detail="Token already spent")
+ elif "Failed to decode token" in error_msg:
+ raise HTTPException(status_code=400, detail="Invalid token format")
+ elif "Invalid token format" in error_msg:
+ raise HTTPException(status_code=400, detail="Invalid token format")
+ elif "Network error" in error_msg:
+ raise HTTPException(
+ status_code=400, detail="Network error during token verification"
+ )
+ else:
+ raise HTTPException(
+ status_code=400, detail=f"Failed to redeem token: {error_msg}"
+ )
+
+ # Zero msats is valid if no error was printed
return {"msats": amount_msats}
@@ -66,8 +122,9 @@ async def refund_wallet_endpoint(
# Perform refund operation first, before modifying balance
if key.refund_address:
+ # refund_balance handles balance update and key deletion
await refund_balance(remaining_balance_msats, key, session)
- result = {"recipient": key.refund_address, "msats": remaining_balance_msats}
+ return {"recipient": key.refund_address, "msats": remaining_balance_msats}
else:
# Convert msats to sats for cashu wallet
remaining_balance_sats = remaining_balance_msats // 1000
@@ -77,17 +134,21 @@ async def refund_wallet_endpoint(
)
# TODO: choose currency and mint based on what user has configured
- token = await wallet().send(remaining_balance_sats)
+ try:
+ token = await wallet().send(remaining_balance_sats)
+ except Exception as e:
+ # Handle mint service errors
+ raise HTTPException(
+ status_code=503, detail=f"Mint service unavailable: {str(e)}"
+ )
- result = {"msats": remaining_balance_msats, "recipient": None, "token": token}
+ # Only for token refunds, we need to manually update balance and delete key
+ key.balance = 0
+ session.add(key)
+ await session.commit()
+ await delete_key_if_zero_balance(key, session)
- # Only after successful refund, zero out the balance
- key.balance = 0
- session.add(key)
- await session.commit()
- await delete_key_if_zero_balance(key, session)
-
- return result
+ return {"msats": remaining_balance_msats, "recipient": None, "token": token}
@wallet_router.api_route(
diff --git a/tests/integration/.env.example b/tests/integration/.env.example
new file mode 100644
index 00000000..444daf22
--- /dev/null
+++ b/tests/integration/.env.example
@@ -0,0 +1,31 @@
+# Integration Test Environment Configuration
+
+# Set to "true" to use real Cashu mint instance instead of mock
+USE_REAL_MINT=false
+
+# URL of the Cashu mint instance (when USE_REAL_MINT=true)
+# For local mint: http://localhost:3338
+# For production mint: https://mint.minibits.cash/Bitcoin
+MINT_URL=http://localhost:3338
+
+# Database configuration (automatically set by tests)
+# DATABASE_URL=sqlite+aiosqlite:///:memory:
+
+# Upstream configuration (for mocking LLM responses)
+UPSTREAM_BASE_URL=https://api.openai.com/v1
+UPSTREAM_API_KEY=test-upstream-key
+
+# Other test configuration
+INTEGRATION_TEST=true
+LOG_LEVEL=DEBUG
+TEST_TIMEOUT=30
+CONCURRENT_TEST_LIMIT=10
+
+# Cashu wallet configuration
+RECEIVE_LN_ADDRESS=test@routstr.com
+REFUND_PROCESSING_INTERVAL=3600
+NSEC=nsec1vl029mgpspedva04g90vltkh6fvh240zqtv9k0t9af8935ke9laqsnlfe5
+COST_PER_REQUEST=10
+MODEL_BASED_PRICING=true
+MINIMUM_PAYOUT=1000
+PAYOUT_INTERVAL=86400
\ No newline at end of file
diff --git a/tests/integration/README.md b/tests/integration/README.md
new file mode 100644
index 00000000..af85bf85
--- /dev/null
+++ b/tests/integration/README.md
@@ -0,0 +1,52 @@
+# Integration Tests
+
+End-to-end tests for API endpoints, Cashu wallet operations, and database interactions.
+
+## Running Tests
+
+```bash
+# All integration tests
+pytest tests/integration/ -v
+
+# Specific test file
+pytest tests/integration/test_wallet_topup.py -v
+
+# Skip slow tests
+pytest tests/integration/ -m "not slow" -v
+```
+
+## Test Infrastructure
+
+**TestmintWallet** - Mock Cashu wallet for generating test tokens
+**DatabaseSnapshot** - Captures database state changes
+**Test Utilities** - Validators for responses, performance, and concurrency
+
+## Real Testmint Setup (Optional)
+
+By default, tests use a mock testmint. For testing against a real instance:
+
+```bash
+./tests/integration/setup_testmint.sh
+export USE_REAL_MINT=true
+export MINT_URL=http://localhost:3338
+pytest tests/integration/ -v
+```
+
+## Writing Tests
+
+```python
+@pytest.mark.integration
+@pytest.mark.asyncio
+async def test_wallet_topup(integration_client, testmint_wallet, db_snapshot):
+ await db_snapshot.capture()
+ token = await testmint_wallet.mint_tokens(1000)
+
+ response = await integration_client.post(
+ "/v1/wallet/topup",
+ params={"cashu_token": token}
+ )
+
+ assert response.status_code == 200
+ diff = await db_snapshot.diff()
+ assert len(diff["api_keys"]["added"]) == 1
+```
\ No newline at end of file
diff --git a/tests/integration/__init__.py b/tests/integration/__init__.py
new file mode 100644
index 00000000..e69de29b
diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py
new file mode 100644
index 00000000..482faa9f
--- /dev/null
+++ b/tests/integration/conftest.py
@@ -0,0 +1,598 @@
+import asyncio
+import json
+import os
+from typing import Any, AsyncGenerator, Callable, Dict, List, Optional, Tuple
+from unittest.mock import AsyncMock, MagicMock, patch
+
+import pytest
+import pytest_asyncio
+from fastapi import FastAPI
+from httpx import AsyncClient
+from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine
+from sqlmodel import select
+
+# Set test environment variables before importing the app
+os.environ.update(
+ {
+ "DATABASE_URL": "sqlite+aiosqlite:///:memory:",
+ "UPSTREAM_BASE_URL": "https://api.openai.com/v1",
+ "UPSTREAM_API_KEY": "test-upstream-key",
+ "MINT": "https://mint.minibits.cash/Bitcoin", # Use real mint URL for tests
+ "RECEIVE_LN_ADDRESS": "test@routstr.com",
+ "REFUND_PROCESSING_INTERVAL": "3600",
+ "NSEC": "nsec1testkey1234567890abcdef",
+ "COST_PER_REQUEST": "10",
+ "MODEL_BASED_PRICING": "true",
+ "MINIMUM_PAYOUT": "1000",
+ "PAYOUT_INTERVAL": "86400",
+ }
+)
+
+from router.db import ApiKey, get_session
+from router.main import app, lifespan
+
+
+class TestmintWallet:
+ """Test wallet that simulates Cashu mint interactions for testing"""
+
+ def __init__(
+ self, mint_url: Optional[str] = None, nsec: Optional[str] = None
+ ) -> None:
+ # Use the configured MINT URL or a local test mint
+ self.mint_url = mint_url or os.environ.get("MINT", "http://localhost:3338")
+ # Use a valid test nsec for testing (this is a well-known test key)
+ self.nsec = (
+ nsec or "nsec1vl029mgpspedva04g90vltkh6fvh240zqtv9k0t9af8935ke9laqsnlfe5"
+ )
+ self.wallet = None
+ self.tokens: List[Dict[str, Any]] = []
+ self.spent_tokens: List[str] = []
+ self.refund_history: List[Dict[str, Any]] = []
+
+ async def init(self) -> None:
+ """Initialize the sixty_nuts wallet"""
+ # In mock mode, we don't actually create a real wallet
+ # This is just a placeholder for the mock implementation
+ self.wallet = None
+
+ async def mint_tokens(self, amount: int) -> str:
+ """Request tokens from testmint - for testing, we simulate this"""
+ # In a real testmint setup, this would request tokens from the mint
+ # For now, we'll create a mock token that the test wallet can "redeem"
+ import base64
+ import secrets
+
+ token_id = secrets.token_hex(16)
+ token_data = {
+ "token": [
+ {
+ "mint": self.mint_url,
+ "proofs": [
+ {
+ "id": token_id,
+ "amount": amount,
+ "secret": secrets.token_hex(32),
+ "C": secrets.token_hex(33),
+ }
+ ],
+ }
+ ],
+ "unit": "sat",
+ "memo": f"Test token {amount} sats",
+ }
+
+ # Encode as Cashu token format
+ token_json = json.dumps(token_data)
+ token_base64 = base64.urlsafe_b64encode(token_json.encode()).decode()
+ cashu_token = f"cashuA{token_base64}"
+
+ self.tokens.append(
+ {"id": token_id, "amount": amount, "token": cashu_token, "spent": False}
+ )
+
+ return cashu_token
+
+ async def redeem_token(self, token: str) -> Tuple[int, str]:
+ """Redeem a Cashu token using the real wallet"""
+ if not self.wallet:
+ await self.init()
+
+ # For testing, simulate the redemption
+ import base64
+
+ if not token.startswith("cashuA"):
+ raise ValueError("Invalid token format")
+
+ try:
+ token_base64 = token[6:] # Remove "cashuA" prefix
+ token_json = base64.urlsafe_b64decode(token_base64).decode()
+ token_data = json.loads(token_json)
+
+ total_amount = 0
+ for mint_tokens in token_data["token"]:
+ for proof in mint_tokens["proofs"]:
+ # Check if token was already spent
+ if proof["id"] in self.spent_tokens:
+ raise ValueError("Token already spent")
+
+ self.spent_tokens.append(proof["id"])
+ total_amount += proof["amount"]
+
+ return total_amount, "test_metadata"
+
+ except Exception as e:
+ raise ValueError(f"Failed to decode token: {str(e)}")
+
+ async def send(self, amount: int) -> str:
+ """Create a token to send (for refunds)"""
+ if not self.wallet:
+ await self.init()
+
+ # For testing, create a refund token
+ return await self.mint_tokens(amount)
+
+ async def send_to_lnurl(self, lnurl: str, amount: int) -> int:
+ """Send to lightning address - simulated for testing"""
+ if not self.wallet:
+ await self.init()
+
+ self.refund_history.append(
+ {
+ "amount": amount,
+ "ln_address": lnurl,
+ "timestamp": asyncio.get_event_loop().time(),
+ }
+ )
+ return amount
+
+ async def get_balance(self) -> int:
+ """Get wallet balance"""
+ if not self.wallet:
+ await self.init()
+
+ # For testing, return a simulated balance
+ return 100000 # 100k sats
+
+
+@pytest_asyncio.fixture
+async def testmint_wallet() -> TestmintWallet:
+ """Fixture for testmint wallet instance"""
+ # Check if we should use real mint
+ mint_url = os.environ.get(
+ "MINT_URL", os.environ.get("MINT", "http://localhost:3338")
+ )
+
+ wallet = TestmintWallet(mint_url=mint_url)
+ await wallet.init()
+ return wallet
+
+
+@pytest_asyncio.fixture
+async def test_database_url(tmp_path: Any) -> str:
+ """Create a temporary SQLite database file for integration tests"""
+ db_file = tmp_path / "test_integration.db"
+ return f"sqlite+aiosqlite:///{db_file}"
+
+
+@pytest_asyncio.fixture
+async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None]:
+ """Create an async engine for integration tests"""
+ engine = create_async_engine(
+ test_database_url,
+ echo=False,
+ future=True,
+ pool_pre_ping=True,
+ pool_size=5,
+ max_overflow=10,
+ )
+
+ # Initialize database schema
+ # Create tables using the engine directly since init_db uses the global engine
+ async with engine.begin() as conn:
+ from sqlmodel import SQLModel
+
+ await conn.run_sync(SQLModel.metadata.create_all)
+
+ yield engine
+
+ # Cleanup
+ await engine.dispose()
+
+
+@pytest_asyncio.fixture
+async def integration_session(
+ integration_engine: Any,
+) -> AsyncGenerator[AsyncSession, None]:
+ """Create a database session for integration tests"""
+ async with AsyncSession(integration_engine, expire_on_commit=False) as session:
+ yield session
+
+
+class DatabaseSnapshot:
+ """Utility to capture and compare database states"""
+
+ def __init__(self, session: AsyncSession) -> None:
+ self.session = session
+ self.snapshot: Optional[Dict[str, List[Dict]]] = None
+
+ async def capture(self) -> Dict[str, List[Dict]]:
+ """Capture current database state"""
+ # Get all API keys with their data
+ result = await self.session.execute(select(ApiKey))
+ api_keys = result.scalars().all()
+
+ snapshot = {
+ "api_keys": [
+ {
+ "hashed_key": key.hashed_key,
+ "balance": key.balance,
+ "total_spent": key.total_spent,
+ "total_requests": key.total_requests,
+ "refund_address": key.refund_address,
+ "key_expiry_time": key.key_expiry_time,
+ }
+ for key in api_keys
+ ]
+ }
+
+ self.snapshot = snapshot
+ return snapshot
+
+ async def diff(
+ self, new_snapshot: Optional[Dict[str, List[Dict]]] = None
+ ) -> Dict[str, Any]:
+ """Calculate differences between snapshots"""
+ if new_snapshot is None:
+ new_snapshot = await self.capture()
+
+ if self.snapshot is None:
+ raise ValueError("No initial snapshot to compare against")
+
+ diff: Dict[str, Dict[str, List[Any]]] = {
+ "api_keys": {"added": [], "removed": [], "modified": []}
+ }
+
+ # Create lookup maps
+ old_keys = {k["hashed_key"]: k for k in self.snapshot["api_keys"]}
+ new_keys = {k["hashed_key"]: k for k in new_snapshot["api_keys"]}
+
+ # Find added keys
+ for key_id in new_keys:
+ if key_id not in old_keys:
+ diff["api_keys"]["added"].append(new_keys[key_id])
+
+ # Find removed keys
+ for key_id in old_keys:
+ if key_id not in new_keys:
+ diff["api_keys"]["removed"].append(old_keys[key_id])
+
+ # Find modified keys
+ for key_id in old_keys:
+ if key_id in new_keys:
+ old = old_keys[key_id]
+ new = new_keys[key_id]
+ changes = {}
+
+ for field in [
+ "balance",
+ "total_spent",
+ "total_requests",
+ "refund_address",
+ "key_expiry_time",
+ ]:
+ if old[field] != new[field]:
+ changes[field] = {
+ "old": old[field],
+ "new": new[field],
+ "delta": new[field] - old[field]
+ if isinstance(new[field], (int, float))
+ else None,
+ }
+
+ if changes:
+ diff["api_keys"]["modified"].append(
+ {"hashed_key": key_id, "changes": changes}
+ )
+
+ return diff
+
+
+@pytest_asyncio.fixture
+async def db_snapshot(integration_session: AsyncSession) -> DatabaseSnapshot:
+ """Database snapshot utility for tracking state changes"""
+ return DatabaseSnapshot(integration_session)
+
+
+@pytest_asyncio.fixture
+async def integration_app(
+ integration_engine: Any,
+ integration_session: AsyncSession,
+ testmint_wallet: TestmintWallet,
+ test_database_url: str,
+) -> AsyncGenerator[FastAPI, None]:
+ """Create FastAPI app instance for integration tests"""
+
+ # Override environment with test database URL
+ os.environ["DATABASE_URL"] = test_database_url
+
+ # Create a new app instance with our lifespan
+ test_app = FastAPI(lifespan=lifespan)
+
+ # Copy all routes from the main app
+ test_app.router = app.router
+
+ # Override the get_session dependency
+ async def override_get_session() -> AsyncGenerator[AsyncSession, None]:
+ yield integration_session
+
+ test_app.dependency_overrides[get_session] = override_get_session
+
+ # Check if we should use real mint
+ use_real_mint = os.environ.get("USE_REAL_MINT", "false").lower() == "true"
+
+ if use_real_mint:
+ # Use real mint with sixty_nuts wallet
+ from .real_testmint import create_real_mint_wallet
+
+ # Create real wallet instance
+ real_wallet = await create_real_mint_wallet()
+
+ with (
+ patch("router.db.engine", integration_engine),
+ patch("router.cashu.wallet_instance", real_wallet.wallet),
+ patch("router.cashu.wallet", lambda: real_wallet.wallet),
+ patch("router.cashu.init_wallet", AsyncMock()),
+ ):
+ yield test_app
+ else:
+ # Use mock testmint wallet (current implementation)
+ with patch("router.db.engine", integration_engine):
+ # Set up the test wallet instance
+ import router.cashu
+
+ original_wallet_instance = router.cashu.wallet_instance
+
+ # Create a wallet adapter that uses our testmint_wallet
+ mock_wallet = AsyncMock()
+ mock_wallet.mint_url = testmint_wallet.mint_url
+ mock_wallet.redeem = testmint_wallet.redeem_token
+ mock_wallet.send = testmint_wallet.send
+ mock_wallet.send_to_lnurl = testmint_wallet.send_to_lnurl
+ mock_wallet.get_balance = testmint_wallet.get_balance
+
+ # Patch the wallet functions to use our test wallet
+ with (
+ patch("router.cashu.wallet") as mock_wallet_func,
+ patch("router.cashu.init_wallet") as mock_init_wallet,
+ ):
+ # Configure to return our test wallet
+ mock_wallet_func.return_value = mock_wallet
+ mock_init_wallet.return_value = None
+
+ # Set the global wallet_instance
+ router.cashu.wallet_instance = mock_wallet
+
+ try:
+ yield test_app
+ finally:
+ # Restore original wallet_instance
+ router.cashu.wallet_instance = original_wallet_instance
+
+
+@pytest_asyncio.fixture
+async def integration_client(
+ integration_app: FastAPI,
+ integration_engine: Any, # Ensure engine is created first
+) -> AsyncGenerator[AsyncClient, None]:
+ """Create an async HTTP client for integration tests"""
+ from httpx import ASGITransport
+
+ async with AsyncClient(
+ transport=ASGITransport(app=integration_app),
+ base_url="http://test",
+ timeout=30.0,
+ ) as client:
+ yield client
+
+
+@pytest_asyncio.fixture
+async def authenticated_client(
+ integration_client: AsyncClient,
+ testmint_wallet: TestmintWallet,
+ integration_session: AsyncSession,
+) -> AsyncClient:
+ """Create an authenticated client with a persistent API key"""
+ # Generate a cashu token
+ test_token = await testmint_wallet.mint_tokens(10000) # 10k sats
+
+ # Use the cashu token as Bearer auth to create an API key
+ integration_client.headers["Authorization"] = f"Bearer {test_token}"
+
+ # Make a request to create the API key (first use of cashu token creates the key)
+ response = await integration_client.get("/v1/wallet/info")
+ assert response.status_code == 200
+ wallet_info = response.json()
+ api_key = wallet_info["api_key"]
+
+ # Now switch to using the persistent API key
+ integration_client.headers["Authorization"] = f"Bearer {api_key}"
+
+ # Store the API key and balance for tests that need it
+ integration_client._test_api_key = api_key # type: ignore
+ integration_client._test_balance = wallet_info["balance"] # type: ignore
+
+ return integration_client
+
+
+@pytest_asyncio.fixture
+async def create_api_key() -> Callable:
+ """Helper to create new API keys for testing"""
+
+ async def _create_key(
+ client: AsyncClient,
+ wallet: TestmintWallet,
+ amount: int = 1000,
+ refund_address: Optional[str] = None,
+ key_expiry_time: Optional[int] = None,
+ ) -> Tuple[str, int]:
+ """Create a new API key and return (api_key, balance)"""
+ # Generate cashu token
+ token = await wallet.mint_tokens(amount)
+
+ # Create headers
+ headers = {"Authorization": f"Bearer {token}"}
+ if refund_address:
+ headers["Refund-LNURL"] = refund_address
+ if key_expiry_time:
+ headers["Key-Expiry-Time"] = str(key_expiry_time)
+
+ # Use the token to create API key
+ response = await client.get("/v1/wallet/info", headers=headers)
+ assert response.status_code == 200
+
+ wallet_info = response.json()
+ return wallet_info["api_key"], wallet_info["balance"]
+
+ return _create_key
+
+
+@pytest.fixture
+def mock_upstream_server() -> Any:
+ """Mock upstream API server responses"""
+ responses: Dict[str, Any] = {}
+
+ class MockResponse:
+ def __init__(
+ self,
+ status_code: int,
+ json_data: Any = None,
+ text_data: Optional[str] = None,
+ ) -> None:
+ self.status_code = status_code
+ self._json_data = json_data
+ self._text_data = text_data
+ self.headers = {"content-type": "application/json"}
+
+ def json(self) -> Any:
+ return self._json_data
+
+ @property
+ def text(self) -> str:
+ return self._text_data or ""
+
+ async def aiter_bytes(
+ self, chunk_size: Optional[int] = None
+ ) -> AsyncGenerator[bytes, None]:
+ """Async iterator for streaming responses"""
+ if self._text_data:
+ yield self._text_data.encode()
+
+ def add_response(method: str, path: str, response: MockResponse) -> None:
+ """Add a mock response for a specific method and path"""
+ responses[f"{method}:{path}"] = response
+
+ def get_response(method: str, path: str) -> MockResponse:
+ """Get mock response for a request"""
+ key = f"{method}:{path}"
+ if key in responses:
+ return responses[key]
+ # Default 404 response
+ return MockResponse(404, {"error": "Not found"})
+
+ mock_server = MagicMock()
+ mock_server.add_response = add_response
+ mock_server.get_response = get_response
+ mock_server.responses = responses
+
+ return mock_server
+
+
+@pytest.fixture
+def integration_env_vars() -> Any:
+ """Fixture to manage integration test environment variables"""
+ original_env = os.environ.copy()
+
+ # Set integration test specific environment variables
+ test_env = {
+ "TESTMINT_URL": "https://testmint.routstr.com",
+ "INTEGRATION_TEST": "true",
+ "LOG_LEVEL": "DEBUG",
+ "DATABASE_POOL_SIZE": "10",
+ "DATABASE_MAX_OVERFLOW": "20",
+ "REQUEST_TIMEOUT": "30",
+ "UPSTREAM_TIMEOUT": "25",
+ }
+
+ os.environ.update(test_env)
+
+ yield test_env
+
+ # Restore original environment
+ os.environ.clear()
+ os.environ.update(original_env)
+
+
+@pytest_asyncio.fixture
+async def background_tasks_controller() -> AsyncGenerator[Any, None]:
+ """Control background tasks during tests"""
+ tasks: List[asyncio.Task] = []
+
+ class TaskController:
+ def __init__(self) -> None:
+ self.paused = False
+ self.cancelled = False
+
+ async def pause(self) -> None:
+ """Pause all background tasks"""
+ self.paused = True
+
+ async def resume(self) -> None:
+ """Resume all background tasks"""
+ self.paused = False
+
+ async def cancel_all(self) -> None:
+ """Cancel all background tasks"""
+ self.cancelled = True
+ for task in tasks:
+ task.cancel()
+
+ controller = TaskController()
+
+ # Patch background task functions to respect controller
+ original_update_pricing: Optional[Callable] = None
+ original_check_refunds: Optional[Callable] = None
+ original_periodic_payout: Optional[Callable] = None
+
+ try:
+ from router.main import check_for_refunds, periodic_payout, update_sats_pricing
+
+ async def controlled_update_pricing() -> None:
+ while not controller.cancelled:
+ if not controller.paused and original_update_pricing:
+ await original_update_pricing()
+ await asyncio.sleep(1)
+
+ async def controlled_check_refunds() -> None:
+ while not controller.cancelled:
+ if not controller.paused and original_check_refunds:
+ await original_check_refunds()
+ await asyncio.sleep(1)
+
+ async def controlled_periodic_payout() -> None:
+ while not controller.cancelled:
+ if not controller.paused and original_periodic_payout:
+ await original_periodic_payout()
+ await asyncio.sleep(1)
+
+ # Store originals and patch
+ original_update_pricing = update_sats_pricing
+ original_check_refunds = check_for_refunds
+ original_periodic_payout = periodic_payout
+
+ except ImportError:
+ pass
+
+ yield controller
+
+ # Cleanup
+ controller.cancelled = True
diff --git a/tests/integration/real_testmint.py b/tests/integration/real_testmint.py
new file mode 100644
index 00000000..2536f84d
--- /dev/null
+++ b/tests/integration/real_testmint.py
@@ -0,0 +1,67 @@
+"""
+Real Cashu mint integration for integration tests.
+
+This module provides a real sixty_nuts Wallet implementation that can be used
+with an actual Cashu mint instance for more thorough integration testing.
+"""
+
+import os
+from typing import Optional, Tuple
+
+from sixty_nuts import Wallet
+
+
+class RealMintWallet:
+ """Real Cashu mint wallet using sixty_nuts library"""
+
+ def __init__(self, mint_url: str, nsec: str):
+ self.mint_url = mint_url
+ self.nsec = nsec
+ self._wallet: Optional[Wallet] = None
+
+ async def init(self) -> None:
+ """Initialize the wallet connection"""
+ if not self._wallet:
+ self._wallet = await Wallet.create(nsec=self.nsec)
+
+ @property
+ def wallet(self) -> Wallet:
+ """Get the wallet instance"""
+ if not self._wallet:
+ raise RuntimeError("Wallet not initialized. Call init() first.")
+ return self._wallet
+
+ async def redeem(self, cashu_token: str) -> Tuple[int, str]:
+ """Redeem a Cashu token"""
+ await self.init()
+ return await self.wallet.redeem(cashu_token)
+
+ async def send(self, amount: int) -> str:
+ """Send amount as Cashu token"""
+ await self.init()
+ return await self.wallet.send(amount)
+
+ async def send_to_lnurl(self, lnurl: str, amount: int) -> int:
+ """Send to lightning address"""
+ await self.init()
+ return await self.wallet.send_to_lnurl(lnurl, amount)
+
+ async def get_balance(self) -> int:
+ """Get wallet balance"""
+ await self.init()
+ return await self.wallet.get_balance()
+
+
+async def create_real_mint_wallet() -> RealMintWallet:
+ """Create a real Cashu mint wallet for integration testing"""
+ mint_url = os.environ.get(
+ "MINT_URL", os.environ.get("MINT", "http://localhost:3338")
+ )
+
+ # Use a valid test nsec (this is a well-known test key)
+ # In production, you would generate a unique key per test run
+ test_nsec = "nsec1vl029mgpspedva04g90vltkh6fvh240zqtv9k0t9af8935ke9laqsnlfe5"
+
+ wallet = RealMintWallet(mint_url=mint_url, nsec=test_nsec)
+ await wallet.init()
+ return wallet
diff --git a/tests/integration/run_performance_tests.py b/tests/integration/run_performance_tests.py
new file mode 100755
index 00000000..5fcec67b
--- /dev/null
+++ b/tests/integration/run_performance_tests.py
@@ -0,0 +1,152 @@
+#!/usr/bin/env python3
+"""
+Performance Testing Runner
+
+This script runs performance tests and generates a detailed report.
+Usage: python tests/integration/run_performance_tests.py
+"""
+
+import asyncio
+import json
+import sys
+from datetime import datetime
+from pathlib import Path
+from typing import Any, Dict, List
+
+# Add project root to path
+sys.path.insert(0, str(Path(__file__).parent.parent.parent))
+
+
+async def run_performance_suite() -> bool:
+ """Run the complete performance test suite"""
+ print("=" * 80)
+ print("ROUTSTR PROXY - PERFORMANCE TEST SUITE")
+ print("=" * 80)
+ print(f"Started at: {datetime.now().isoformat()}")
+ print()
+
+ # Performance test commands
+ test_suites = [
+ {
+ "name": "Baseline Performance Metrics",
+ "cmd": "pytest tests/integration/test_performance_load.py::TestPerformanceBaseline -v -s",
+ },
+ {
+ "name": "Load Testing - 100 Concurrent Users",
+ "cmd": "pytest tests/integration/test_performance_load.py::TestLoadScenarios::test_concurrent_users_100 -v -s",
+ },
+ {
+ "name": "Sustained Load - 1000 RPM",
+ "cmd": "pytest tests/integration/test_performance_load.py::TestLoadScenarios::test_sustained_load_1000_rpm -v -s",
+ },
+ {
+ "name": "Memory Leak Detection",
+ "cmd": "pytest tests/integration/test_performance_load.py::TestMemoryLeaks -v -s",
+ },
+ {
+ "name": "Performance Regression Tests",
+ "cmd": "pytest tests/integration/test_performance_load.py::TestPerformanceRegression -v -s",
+ },
+ ]
+
+ results: List[Dict[str, Any]] = []
+
+ for suite in test_suites:
+ print(f"\n{'=' * 60}")
+ print(f"Running: {suite['name']}")
+ print(f"{'=' * 60}")
+
+ start_time = datetime.now()
+
+ # Run the test
+ proc = await asyncio.create_subprocess_shell(
+ suite["cmd"], stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE
+ )
+
+ stdout, stderr = await proc.communicate()
+
+ end_time = datetime.now()
+ duration = (end_time - start_time).total_seconds()
+
+ result = {
+ "name": suite["name"],
+ "success": proc.returncode == 0,
+ "duration": duration,
+ "start_time": start_time.isoformat(),
+ "end_time": end_time.isoformat(),
+ }
+
+ if proc.returncode == 0:
+ print(f"PASSED: {suite['name']} ({duration:.2f}s)")
+ else:
+ print(f"FAILED: {suite['name']} ({duration:.2f}s)")
+ if stderr:
+ print(f"Error: {stderr.decode()}")
+
+ results.append(result)
+
+ # Generate report
+ print("\n" + "=" * 80)
+ print("PERFORMANCE TEST SUMMARY")
+ print("=" * 80)
+
+ total_tests = len(results)
+ passed_tests = sum(1 for r in results if r["success"])
+ failed_tests = total_tests - passed_tests
+
+ print(f"Total Tests: {total_tests}")
+ print(f"Passed: {passed_tests}")
+ print(f"Failed: {failed_tests}")
+ print(f"Success Rate: {(passed_tests / total_tests) * 100:.1f}%")
+
+ # Save report
+ report_dir = Path("tests/integration/performance_reports")
+ report_dir.mkdir(exist_ok=True)
+
+ report_file = (
+ report_dir
+ / f"performance_report_{datetime.now().strftime('%Y%m%d_%H%M%S')}.json"
+ )
+
+ report_data = {
+ "timestamp": datetime.now().isoformat(),
+ "summary": {
+ "total": total_tests,
+ "passed": passed_tests,
+ "failed": failed_tests,
+ "success_rate": passed_tests / total_tests,
+ },
+ "results": results,
+ }
+
+ with open(report_file, "w") as f:
+ json.dump(report_data, f, indent=2)
+
+ print(f"\nDetailed report saved to: {report_file}")
+
+ return passed_tests == total_tests
+
+
+async def main() -> None:
+ """Main entry point"""
+ # Check if proxy server is running
+ import httpx
+
+ try:
+ async with httpx.AsyncClient() as client:
+ response = await client.get("http://localhost:8000/")
+ if response.status_code != 200:
+ print("WARNING: Proxy server may not be running properly")
+ except Exception:
+ print("ERROR: Proxy server is not running!")
+ print("Please start the server with: uvicorn router.main:app")
+ sys.exit(1)
+
+ # Run performance tests
+ success = await run_performance_suite()
+
+ sys.exit(0 if success else 1)
+
+
+if __name__ == "__main__":
+ asyncio.run(main())
diff --git a/tests/integration/setup_cashu_mint.sh b/tests/integration/setup_cashu_mint.sh
new file mode 100755
index 00000000..58ed8908
--- /dev/null
+++ b/tests/integration/setup_cashu_mint.sh
@@ -0,0 +1,56 @@
+#!/bin/bash
+
+# Script to set up a local Cashu mint instance for integration testing
+
+echo "Setting up local Cashu mint instance..."
+
+# Check if Docker is installed
+if ! command -v docker &> /dev/null; then
+ echo "Error: Docker is not installed. Please install Docker first."
+ exit 1
+fi
+
+# Stop any existing mint container
+echo "Stopping any existing Cashu mint container..."
+docker stop cashu-mint-test 2>/dev/null || true
+docker rm cashu-mint-test 2>/dev/null || true
+
+# Start Cashu mint container
+echo "Starting Cashu mint container..."
+docker run -d \
+ --name cashu-mint-test \
+ -p 3338:3338 \
+ -e MINT_BACKEND_BOLT11_SAT=FakeWallet \
+ -e MINT_LISTEN_HOST=0.0.0.0 \
+ -e MINT_LISTEN_PORT=3338 \
+ -e MINT_PRIVATE_KEY=supersecretprivatekey \
+ cashubtc/nutshell:latest \
+ mint
+
+# Wait for mint to be ready
+echo "Waiting for Cashu mint to be ready..."
+for i in {1..30}; do
+ if curl -f http://localhost:3338/v1/info >/dev/null 2>&1; then
+ echo "Cashu mint is ready!"
+ break
+ fi
+ if [ $i -eq 30 ]; then
+ echo "Error: Cashu mint failed to start within 30 seconds"
+ docker logs cashu-mint-test
+ exit 1
+ fi
+ sleep 1
+done
+
+# Display connection info
+echo ""
+echo "Cashu mint is running at: http://localhost:3338"
+echo ""
+echo "To run integration tests with real Cashu mint:"
+echo " export USE_REAL_MINT=true"
+echo " export MINT_URL=http://localhost:3338"
+echo " pytest tests/integration/ -v"
+echo ""
+echo "To stop Cashu mint:"
+echo " docker stop cashu-mint-test"
+echo " docker rm cashu-mint-test"
\ No newline at end of file
diff --git a/tests/integration/test_background_tasks.py b/tests/integration/test_background_tasks.py
new file mode 100644
index 00000000..90ee0cd2
--- /dev/null
+++ b/tests/integration/test_background_tasks.py
@@ -0,0 +1,747 @@
+"""Integration tests for background tasks"""
+
+import asyncio
+import os
+import time
+from datetime import datetime, timedelta
+from typing import Any, Coroutine, List
+from unittest.mock import AsyncMock, patch
+
+import pytest
+
+from router.cashu import check_for_refunds, periodic_payout
+from router.db import ApiKey
+from router.models import MODELS, Model, Pricing, update_sats_pricing
+
+
+@pytest.mark.asyncio
+class TestPricingUpdateTask:
+ """Test the pricing update background task"""
+
+ async def test_updates_model_prices_periodically(self) -> None:
+ """Test that update_sats_pricing updates all model prices based on BTC/USD rate"""
+ # Mock the price fetch function
+ mock_sats_usd = 0.00002 # 1 sat = $0.00002 (BTC at $50,000)
+
+ with patch(
+ "router.models.sats_usd_ask_price", AsyncMock(return_value=mock_sats_usd)
+ ):
+ # Create a test model
+ test_model = Model( # type: ignore[arg-type]
+ id="test-model",
+ name="Test Model",
+ created=1234567890,
+ description="Test",
+ context_length=4096,
+ architecture={ # type: ignore[arg-type]
+ "modality": "text",
+ "input_modalities": ["text"],
+ "output_modalities": ["text"],
+ "tokenizer": "test",
+ "instruct_type": None,
+ },
+ pricing=Pricing(
+ prompt=0.001, # $0.001 per token
+ completion=0.002,
+ request=0.0,
+ image=0.0,
+ web_search=0.0,
+ internal_reasoning=0.0,
+ max_cost=0.0,
+ ),
+ top_provider={ # type: ignore[arg-type]
+ "context_length": 4096,
+ "max_completion_tokens": 1024,
+ "is_moderated": False,
+ },
+ )
+
+ # Add test model to MODELS list
+ original_models = MODELS.copy()
+ MODELS.clear()
+ MODELS.append(test_model)
+
+ try:
+ # Run the pricing update task once
+ task = asyncio.create_task(update_sats_pricing())
+ await asyncio.sleep(0.1) # Let it run one iteration
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+
+ # Verify sats pricing was calculated correctly
+ assert test_model.sats_pricing is not None
+ assert test_model.sats_pricing.prompt == pytest.approx(
+ 0.001 / mock_sats_usd
+ )
+ assert test_model.sats_pricing.completion == pytest.approx(
+ 0.002 / mock_sats_usd
+ )
+
+ # Verify max_cost calculation
+ expected_max_cost = (
+ 4096 * test_model.sats_pricing.prompt
+ + 1024 * test_model.sats_pricing.completion
+ )
+ assert test_model.sats_pricing.max_cost == pytest.approx(
+ expected_max_cost
+ )
+
+ finally:
+ # Restore original models
+ MODELS.clear()
+ MODELS.extend(original_models)
+
+ async def test_handles_provider_api_failures(self) -> None:
+ """Test that pricing update continues running even if price API fails"""
+ call_count = 0
+
+ async def mock_price_func() -> float:
+ nonlocal call_count
+ call_count += 1
+ if call_count == 1:
+ raise Exception("Price API error")
+ return 0.00002
+
+ with patch("router.models.sats_usd_ask_price", mock_price_func):
+ # Run the task
+ task = asyncio.create_task(update_sats_pricing())
+ await asyncio.sleep(15) # Let it run for >10 seconds (one retry)
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+
+ # Verify it retried after the error
+ assert call_count >= 2
+
+ async def test_database_updates_are_atomic(self) -> None:
+ """Test that model price updates don't interfere with concurrent operations"""
+ # This test verifies the pricing updates are in-memory only
+ # and don't affect database operations
+
+ test_model = Model( # type: ignore[arg-type]
+ id="test-atomic",
+ name="Test Atomic",
+ created=1234567890,
+ description="Test",
+ context_length=4096,
+ architecture={ # type: ignore[arg-type]
+ "modality": "text",
+ "input_modalities": ["text"],
+ "output_modalities": ["text"],
+ "tokenizer": "test",
+ "instruct_type": None,
+ },
+ pricing=Pricing(
+ prompt=0.001,
+ completion=0.002,
+ request=0.0,
+ image=0.0,
+ web_search=0.0,
+ internal_reasoning=0.0,
+ max_cost=0.0,
+ ),
+ )
+
+ original_models = MODELS.copy()
+ MODELS.clear()
+ MODELS.append(test_model)
+
+ try:
+ with patch(
+ "router.models.sats_usd_ask_price", AsyncMock(return_value=0.00002)
+ ):
+ # Start the pricing task
+ task = asyncio.create_task(update_sats_pricing())
+
+ # Simulate concurrent access to the model
+ results = []
+
+ async def access_model() -> None:
+ await asyncio.sleep(0.05) # Small delay
+ results.append(test_model.sats_pricing)
+
+ # Run multiple concurrent accesses during pricing update
+ await asyncio.gather(*[access_model() for _ in range(10)])
+
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+
+ # All accesses should see consistent state
+ assert all(r is not None for r in results)
+
+ finally:
+ MODELS.clear()
+ MODELS.extend(original_models)
+
+
+@pytest.mark.asyncio
+class TestRefundCheckTask:
+ """Test the refund check background task"""
+
+ async def test_processes_pending_refunds(
+ self, integration_session: Any, testmint_wallet: Any, db_snapshot: Any
+ ) -> None:
+ """Test that expired keys with balance and refund address are refunded"""
+ # Create an expired API key with balance
+ expired_key = ApiKey(
+ hashed_key="expired_test_key",
+ balance=5000, # 5 sats in msats
+ refund_address="lnurl1test",
+ key_expiry_time=int(time.time()) - 3600, # Expired 1 hour ago
+ created_at=datetime.utcnow() - timedelta(days=1),
+ )
+ integration_session.add(expired_key)
+ await integration_session.commit()
+
+ # Mock the wallet send_to_lnurl method and get_session
+ with (
+ patch("router.cashu.wallet") as mock_wallet,
+ patch("router.cashu.get_session") as mock_get_session,
+ ):
+ mock_wallet_instance = AsyncMock()
+ mock_wallet_instance.send_to_lnurl = AsyncMock(return_value=5)
+ mock_wallet.return_value = mock_wallet_instance
+
+ # Make get_session return our integration session
+ async def get_test_session() -> Any:
+ yield integration_session
+
+ mock_get_session.side_effect = get_test_session
+
+ # Take initial snapshot
+ await db_snapshot.capture()
+
+ # Run refund check once
+ original_interval = os.environ.get("REFUND_PROCESSING_INTERVAL", "3600")
+ os.environ["REFUND_PROCESSING_INTERVAL"] = (
+ "0.1" # Fast interval for testing
+ )
+
+ task = asyncio.create_task(check_for_refunds())
+ await asyncio.sleep(0.5) # Let it run one cycle
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+ finally:
+ os.environ["REFUND_PROCESSING_INTERVAL"] = original_interval
+
+ # Verify refund was processed
+ mock_wallet_instance.send_to_lnurl.assert_called_once_with(
+ "lnurl1test", amount=5
+ )
+
+ # Check database state
+ db_diff = await db_snapshot.diff()
+ assert len(db_diff["api_keys"]["modified"]) == 1
+ modified_key = db_diff["api_keys"]["modified"][0]
+ assert modified_key["changes"]["balance"]["new"] == 0
+ assert modified_key["changes"]["balance"]["delta"] == -5000
+
+ async def test_handles_mint_communication_errors(
+ self, integration_session: Any
+ ) -> None:
+ """Test that refund check continues after mint errors"""
+ # Create multiple expired keys
+ for i in range(3):
+ key = ApiKey(
+ hashed_key=f"expired_key_{i}",
+ balance=1000 * (i + 1),
+ refund_address=f"lnurl{i}",
+ key_expiry_time=int(time.time()) - 3600,
+ created_at=datetime.utcnow(),
+ )
+ integration_session.add(key)
+ await integration_session.commit()
+
+ refund_count = 0
+
+ async def mock_send_to_lnurl(address: str, amount: int) -> int:
+ nonlocal refund_count
+ refund_count += 1
+ if refund_count == 2:
+ raise Exception("Mint communication error")
+ return amount
+
+ with (
+ patch("router.cashu.wallet") as mock_wallet,
+ patch("router.cashu.get_session") as mock_get_session,
+ ):
+ mock_wallet_instance = AsyncMock()
+ mock_wallet_instance.send_to_lnurl = mock_send_to_lnurl
+ mock_wallet.return_value = mock_wallet_instance
+
+ # Make get_session return our integration session
+ async def get_test_session() -> Any:
+ yield integration_session
+
+ mock_get_session.side_effect = get_test_session
+
+ # Run refund check
+ original_interval = os.environ.get("REFUND_PROCESSING_INTERVAL", "3600")
+ os.environ["REFUND_PROCESSING_INTERVAL"] = "0.1"
+
+ task = asyncio.create_task(check_for_refunds())
+ await asyncio.sleep(0.5)
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+ finally:
+ os.environ["REFUND_PROCESSING_INTERVAL"] = original_interval
+
+ # Should have attempted all refunds despite one failure
+ assert refund_count == 3
+
+ async def test_updates_refund_status_correctly(
+ self, integration_session: Any, db_snapshot: Any
+ ) -> None:
+ """Test that refund status and key deletion work correctly"""
+ # Create keys with different states
+ keys_data = [
+ # Should be refunded and deleted (zero balance after refund)
+ {
+ "hashed_key": "delete_me",
+ "balance": 1000,
+ "refund_address": "lnurl1",
+ "expired": True,
+ },
+ # Should keep (not expired)
+ {
+ "hashed_key": "keep_not_expired",
+ "balance": 2000,
+ "refund_address": "lnurl2",
+ "expired": False,
+ },
+ # Should keep (no refund address)
+ {
+ "hashed_key": "keep_no_address",
+ "balance": 3000,
+ "refund_address": None,
+ "expired": True,
+ },
+ # Already zero balance
+ {
+ "hashed_key": "zero_balance",
+ "balance": 0,
+ "refund_address": "lnurl3",
+ "expired": True,
+ },
+ ]
+
+ current_time = int(time.time())
+ for data in keys_data:
+ key = ApiKey(
+ hashed_key=data["hashed_key"],
+ balance=data["balance"],
+ refund_address=data["refund_address"],
+ key_expiry_time=current_time - 3600
+ if data["expired"]
+ else current_time + 3600,
+ created_at=datetime.utcnow(),
+ )
+ integration_session.add(key)
+ await integration_session.commit()
+
+ with (
+ patch("router.cashu.wallet") as mock_wallet,
+ patch("router.cashu.get_session") as mock_get_session,
+ ):
+ mock_wallet_instance = AsyncMock()
+ mock_wallet_instance.send_to_lnurl = AsyncMock(return_value=1)
+ mock_wallet.return_value = mock_wallet_instance
+
+ # Make get_session return our integration session
+ async def get_test_session() -> Any:
+ yield integration_session
+
+ mock_get_session.side_effect = get_test_session
+
+ await db_snapshot.capture()
+
+ # Run refund check
+ original_interval = os.environ.get("REFUND_PROCESSING_INTERVAL", "3600")
+ os.environ["REFUND_PROCESSING_INTERVAL"] = "0.1"
+
+ task = asyncio.create_task(check_for_refunds())
+ await asyncio.sleep(0.5)
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+ finally:
+ os.environ["REFUND_PROCESSING_INTERVAL"] = original_interval
+
+ # Verify correct keys were processed
+ assert mock_wallet_instance.send_to_lnurl.call_count == 1
+ mock_wallet_instance.send_to_lnurl.assert_called_with("lnurl1", amount=1)
+
+ # Check final state
+ from sqlalchemy import select as sa_select
+
+ result = await integration_session.execute(sa_select(ApiKey))
+ remaining_keys_list = result.scalars().all()
+ remaining_ids = [k.hashed_key for k in remaining_keys_list]
+
+ assert "delete_me" not in remaining_ids # Deleted after refund
+ assert "keep_not_expired" in remaining_ids
+ assert "keep_no_address" in remaining_ids
+ assert (
+ "zero_balance" not in remaining_ids
+ ) # Auto-deleted due to zero balance
+
+ async def test_refund_check_disabled(self) -> None:
+ """Test that refund check can be disabled by setting interval to 0"""
+ # Set refund interval to 0 to disable
+ original_interval = os.environ.get("REFUND_PROCESSING_INTERVAL", "3600")
+ os.environ["REFUND_PROCESSING_INTERVAL"] = "0"
+
+ try:
+ # Task should exit immediately
+ task = asyncio.create_task(check_for_refunds())
+ await task # Should complete without hanging
+
+ # Task should have exited cleanly
+ assert task.done()
+ finally:
+ os.environ["REFUND_PROCESSING_INTERVAL"] = original_interval
+
+
+@pytest.mark.asyncio
+class TestPeriodicPayoutTask:
+ """Test the periodic payout background task"""
+
+ async def test_executes_at_configured_intervals(self) -> None:
+ """Test that payout task runs at the configured interval"""
+ call_count = 0
+
+ async def mock_pay_out() -> None:
+ nonlocal call_count
+ call_count += 1
+
+ with patch("router.cashu.pay_out", mock_pay_out):
+ # Set a short interval for testing
+ original_interval = os.environ.get("PAYOUT_INTERVAL", "300")
+ os.environ["PAYOUT_INTERVAL"] = "0.2" # 200ms
+
+ task = asyncio.create_task(periodic_payout())
+ await asyncio.sleep(0.7) # Should run ~3 times
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+ finally:
+ os.environ["PAYOUT_INTERVAL"] = original_interval
+
+ assert 2 <= call_count <= 4 # Allow some timing variance
+
+ async def test_calculates_payouts_accurately(
+ self, integration_session: Any
+ ) -> None:
+ """Test that payouts are calculated correctly based on revenue"""
+ # Create test API keys with various balances
+ total_user_balance = 0
+ for i in range(5):
+ balance = 10000 * (i + 1) # 10, 20, 30, 40, 50 sats
+ total_user_balance += balance
+ key = ApiKey(
+ hashed_key=f"user_key_{i}",
+ balance=balance,
+ created_at=datetime.utcnow(),
+ )
+ integration_session.add(key)
+ await integration_session.commit()
+
+ # Mock wallet balance higher than user balances (indicating revenue)
+ wallet_balance = 200000 # 200 sats total
+ expected_revenue = wallet_balance - total_user_balance # 50 sats revenue
+
+ with patch("router.cashu.wallet") as mock_wallet:
+ mock_wallet_instance = AsyncMock()
+ mock_wallet_instance.balance = AsyncMock(return_value=wallet_balance)
+ mock_wallet_instance.send_to_lnurl = AsyncMock(return_value=None)
+ mock_wallet.return_value = mock_wallet_instance
+
+ # Mock environment variables
+ with patch.dict(
+ os.environ,
+ {
+ "MINIMUM_PAYOUT": "10", # 10 sats minimum
+ "RECEIVE_LN_ADDRESS": "owner@test.com",
+ "DEV_LN_ADDRESS": "dev@test.com",
+ },
+ ):
+ # Call pay_out directly
+ from router.cashu import pay_out
+
+ await pay_out()
+
+ # Verify payouts were sent correctly
+ assert mock_wallet_instance.send_to_lnurl.call_count == 2
+
+ # Check amounts (97.9% to owner, 2.1% to dev)
+ calls = mock_wallet_instance.send_to_lnurl.call_args_list
+ owner_call = next(c for c in calls if c[0][0] == "owner@test.com")
+ dev_call = next(c for c in calls if c[0][0] == "dev@test.com")
+
+ owner_amount = owner_call[0][1]
+ dev_amount = dev_call[0][1]
+
+ assert owner_amount == int(expected_revenue * 0.979)
+ assert dev_amount == int(expected_revenue * 0.021)
+ assert owner_amount + dev_amount == expected_revenue
+
+ async def test_transaction_logging_complete(
+ self, integration_session: Any, capfd: Any
+ ) -> None:
+ """Test that payout transactions are properly logged"""
+ # Create a simple scenario
+ key = ApiKey(
+ hashed_key="single_user",
+ balance=50000, # 50 sats
+ created_at=datetime.utcnow(),
+ )
+ integration_session.add(key)
+ await integration_session.commit()
+
+ with patch("router.cashu.wallet") as mock_wallet:
+ mock_wallet_instance = AsyncMock()
+ mock_wallet_instance.balance = AsyncMock(
+ return_value=100000
+ ) # 100 sats total
+ mock_wallet_instance.send_to_lnurl = AsyncMock(return_value=None)
+ mock_wallet.return_value = mock_wallet_instance
+
+ with patch.dict(
+ os.environ,
+ {
+ "MINIMUM_PAYOUT": "10",
+ "RECEIVE_LN_ADDRESS": "owner@test.com",
+ "DEV_LN_ADDRESS": "dev@test.com",
+ },
+ ):
+ from router.cashu import pay_out
+
+ await pay_out()
+
+ # Check that logging occurred
+ captured = capfd.readouterr()
+ assert "Revenue:" in captured.out
+ assert "Owner's draw:" in captured.out
+ assert "Developer's donation:" in captured.out
+
+ async def test_minimum_payout_threshold(self, integration_session: Any) -> None:
+ """Test that payouts only occur when revenue exceeds minimum threshold"""
+ # Create scenario with low revenue
+ key = ApiKey(
+ hashed_key="low_revenue_user",
+ balance=95000, # 95 sats
+ created_at=datetime.utcnow(),
+ )
+ integration_session.add(key)
+ await integration_session.commit()
+
+ with patch("router.cashu.wallet") as mock_wallet:
+ mock_wallet_instance = AsyncMock()
+ mock_wallet_instance.balance = AsyncMock(
+ return_value=96000
+ ) # Only 1 sat revenue
+ mock_wallet_instance.send_to_lnurl = AsyncMock(return_value=None)
+ mock_wallet.return_value = mock_wallet_instance
+
+ with patch.dict(os.environ, {"MINIMUM_PAYOUT": "10"}): # 10 sats minimum
+ from router.cashu import pay_out
+
+ await pay_out()
+
+ # No payouts should have been sent
+ mock_wallet_instance.send_to_lnurl.assert_not_called()
+
+
+@pytest.mark.asyncio
+class TestTaskInteractions:
+ """Test interactions between background tasks"""
+
+ async def test_tasks_dont_interfere_with_each_other(self) -> None:
+ """Test that all tasks can run concurrently without issues"""
+ # Mock all external dependencies
+ with (
+ patch("router.models.sats_usd_ask_price", AsyncMock(return_value=0.00002)),
+ patch("router.cashu.wallet") as mock_wallet,
+ patch("router.cashu.pay_out", AsyncMock()),
+ ):
+ mock_wallet_instance = AsyncMock()
+ mock_wallet_instance.send_to_lnurl = AsyncMock(return_value=1)
+ mock_wallet.return_value = mock_wallet_instance
+
+ # Start all tasks
+ tasks = []
+ try:
+ # Pricing task
+ pricing_task = asyncio.create_task(update_sats_pricing())
+ tasks.append(pricing_task)
+
+ # Refund task (disabled to avoid interference)
+ os.environ["REFUND_PROCESSING_INTERVAL"] = "0"
+ refund_task = asyncio.create_task(check_for_refunds())
+ tasks.append(refund_task)
+
+ # Payout task
+ payout_task = asyncio.create_task(periodic_payout())
+ tasks.append(payout_task)
+
+ # Let them run concurrently
+ await asyncio.sleep(0.5)
+
+ # All tasks should still be running (except refund which exits immediately)
+ assert not pricing_task.done()
+ assert refund_task.done() # Should exit immediately when disabled
+ assert not payout_task.done()
+
+ finally:
+ # Clean up
+ for task in tasks:
+ if not task.done():
+ task.cancel()
+ await asyncio.gather(*tasks, return_exceptions=True)
+
+ async def test_api_requests_work_during_task_execution(
+ self, integration_client: Any
+ ) -> None:
+ """Test that API endpoints remain responsive during background task execution"""
+ # Start a mock long-running task
+ processing = asyncio.Event()
+
+ async def slow_task() -> None:
+ processing.set()
+ await asyncio.sleep(2) # Simulate long operation
+
+ with patch("router.models.sats_usd_ask_price", slow_task):
+ # Start the pricing task
+ task = asyncio.create_task(update_sats_pricing())
+
+ # Wait for task to start processing
+ await processing.wait()
+
+ # API should still be responsive
+ response = await integration_client.get("/")
+ assert response.status_code == 200
+
+ # Models endpoint should work
+ response = await integration_client.get("/v1/models")
+ assert response.status_code == 200
+
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+
+ async def test_database_locking_handled_properly(
+ self, integration_session: Any
+ ) -> None:
+ """Test that database operations don't deadlock during concurrent task execution"""
+ # Create test data
+ for i in range(10):
+ key = ApiKey(
+ hashed_key=f"concurrent_key_{i}",
+ balance=1000 * i,
+ refund_address=f"lnurl{i}" if i % 2 == 0 else None,
+ key_expiry_time=int(time.time()) - 3600
+ if i % 3 == 0
+ else int(time.time()) + 3600,
+ created_at=datetime.utcnow(),
+ )
+ integration_session.add(key)
+ await integration_session.commit()
+
+ # Simulate concurrent database operations
+ async def read_operation() -> int:
+ from sqlalchemy import select as sa_select
+
+ result = await integration_session.execute(sa_select(ApiKey))
+ return len(result.scalars().all())
+
+ async def write_operation(key_id: int) -> None:
+ from sqlalchemy import select as sa_select
+
+ stmt = sa_select(ApiKey).where(
+ ApiKey.hashed_key == f"concurrent_key_{key_id}" # type: ignore[arg-type]
+ )
+ result = await integration_session.execute(stmt)
+ key = result.scalar_one_or_none()
+ if key:
+ key.balance += 100
+ await integration_session.commit()
+
+ # Run multiple operations concurrently
+ tasks: List[Coroutine[Any, Any, Any]] = []
+ for _ in range(5):
+ tasks.append(read_operation()) # type: ignore[arg-type]
+ for i in range(5):
+ tasks.append(write_operation(i)) # type: ignore[arg-type]
+
+ # All operations should complete without deadlock
+ results = await asyncio.gather(*tasks, return_exceptions=True)
+
+ # Check no exceptions occurred
+ exceptions = [r for r in results if isinstance(r, Exception)]
+ assert len(exceptions) == 0
+
+ async def test_graceful_shutdown(self) -> None:
+ """Test that all tasks shut down cleanly when cancelled"""
+ shutdown_messages = []
+
+ async def task_with_cleanup(name: str) -> None:
+ try:
+ while True:
+ await asyncio.sleep(0.1)
+ except asyncio.CancelledError:
+ shutdown_messages.append(f"{name} shutting down")
+ raise
+
+ # Patch the actual task functions
+ with (
+ patch(
+ "router.models.update_sats_pricing",
+ lambda: task_with_cleanup("pricing"),
+ ),
+ patch(
+ "router.cashu.check_for_refunds", lambda: task_with_cleanup("refund")
+ ),
+ patch("router.cashu.periodic_payout", lambda: task_with_cleanup("payout")),
+ ):
+ # Start all tasks
+ tasks = [
+ asyncio.create_task(update_sats_pricing()),
+ asyncio.create_task(check_for_refunds()),
+ asyncio.create_task(periodic_payout()),
+ ]
+
+ # Let them start
+ await asyncio.sleep(0.2)
+
+ # Cancel all tasks
+ for task in tasks:
+ task.cancel()
+
+ # Wait for cleanup
+ await asyncio.gather(*tasks, return_exceptions=True)
+
+ # Verify all tasks shut down properly
+ assert len(shutdown_messages) == 3
+ assert "pricing shutting down" in shutdown_messages
+ assert "refund shutting down" in shutdown_messages
+ assert "payout shutting down" in shutdown_messages
diff --git a/tests/integration/test_database_consistency.py b/tests/integration/test_database_consistency.py
new file mode 100644
index 00000000..046ba006
--- /dev/null
+++ b/tests/integration/test_database_consistency.py
@@ -0,0 +1,618 @@
+"""Comprehensive database consistency tests"""
+
+import asyncio
+import time
+from typing import Any, Dict, List
+from unittest.mock import AsyncMock, MagicMock, patch
+
+import pytest
+from httpx import AsyncClient, Response
+from sqlalchemy.exc import IntegrityError
+from sqlalchemy.ext.asyncio import AsyncSession
+from sqlmodel import select
+
+from router.db import ApiKey
+
+
+class TestTransactionAtomicity:
+ """Test transaction atomicity across all database operations"""
+
+ @pytest.mark.asyncio
+ async def test_balance_update_atomicity(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ db_snapshot: Any,
+ ) -> None:
+ """Test that balance updates are atomic and rolled back on failure"""
+ # Get initial balance
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ # API key format is "sk-{hashed_key}"
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ api_key = result.scalar_one()
+ initial_balance = api_key.balance
+
+ # Test database atomicity by simulating a failed transaction
+ # Create a new session for isolated transaction
+ from sqlalchemy.ext.asyncio import AsyncSession
+
+ async with AsyncSession(integration_session.bind) as test_session:
+ try:
+ # Get api key in new session
+ result = await test_session.execute(
+ select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ )
+ test_api_key = result.scalar_one()
+
+ # Update balance
+ test_api_key.balance -= 1000
+ await test_session.flush() # Apply changes but don't commit
+
+ # Simulate an error that would cause rollback
+ raise Exception("Simulated error after balance update")
+ except Exception:
+ await test_session.rollback()
+
+ # Verify balance wasn't changed in main session
+ await integration_session.refresh(api_key)
+ assert api_key.balance == initial_balance
+
+ # Test with concurrent modifications
+ await db_snapshot.capture()
+
+ # Try to update in a transaction that will fail
+ from sqlalchemy import update
+
+ try:
+ await integration_session.execute(
+ update(ApiKey)
+ .where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ .values(balance=ApiKey.balance - 1000)
+ )
+ # Force a constraint violation or error
+ await integration_session.execute(
+ update(ApiKey)
+ .where(ApiKey.hashed_key == "non_existent_key") # type: ignore[arg-type]
+ .values(balance=-1) # This should fail
+ )
+ await integration_session.commit()
+ except Exception:
+ await integration_session.rollback()
+
+ # Verify no changes were persisted
+ diff = await db_snapshot.diff()
+ assert len(diff["api_keys"]["added"]) == 0
+ assert len(diff["api_keys"]["modified"]) == 0
+
+ @pytest.mark.asyncio
+ async def test_topup_rollback_on_failure(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ db_snapshot: Any,
+ ) -> None:
+ """Test that failed top-ups don't leave partial database state"""
+ # Get initial state
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ # API key format is "sk-{hashed_key}"
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ api_key = result.scalar_one()
+ initial_balance = api_key.balance
+
+ # Mock wallet to fail after token validation
+ with patch("router.cashu.wallet") as mock_wallet_func:
+ mock_proof = MagicMock()
+ mock_proof.amount = 1000
+ mock_wallet = AsyncMock()
+ mock_wallet.deserialize_token = AsyncMock(return_value=[mock_proof])
+ mock_wallet.redeem = AsyncMock(
+ side_effect=Exception("Network error during redemption")
+ )
+ mock_wallet_func.return_value = mock_wallet
+
+ # Attempt top-up
+ response = await authenticated_client.post(
+ "/v1/wallet/topup", params={"cashu_token": "cashuAey..."}
+ )
+
+ # The mock returns 400 for invalid tokens
+ assert response.status_code in [400, 500]
+
+ # Verify no balance change
+ await integration_session.refresh(api_key)
+ assert api_key.balance == initial_balance
+
+ # Verify clean database state
+ diff = await db_snapshot.diff()
+ assert len(diff["api_keys"]["added"]) == 0
+ assert len(diff["api_keys"]["modified"]) == 0
+
+ @pytest.mark.asyncio
+ async def test_concurrent_balance_updates(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test atomic balance updates under concurrent operations"""
+ # Get API key info
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ # API key format is "sk-{hashed_key}"
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ # Set a known balance
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ api_key = result.scalar_one()
+ api_key.balance = 10000
+ await integration_session.commit()
+
+ # Simulate concurrent balance updates through direct database operations
+ async def update_balance(session: AsyncSession, amount: int) -> bool:
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await session.execute(stmt)
+ key = result.scalar_one()
+ key.balance -= amount
+ key.total_spent += amount
+ key.total_requests += 1
+ try:
+ await session.commit()
+ return True
+ except Exception:
+ await session.rollback()
+ return False
+
+ # Run concurrent balance updates
+ tasks = []
+ deduction_amounts = [100, 200, 300, 400, 500]
+
+ for amount in deduction_amounts:
+ # Create a new session for each concurrent operation
+ async with AsyncSession(integration_session.bind) as session:
+ task = update_balance(session, amount)
+ tasks.append(task)
+
+ await asyncio.gather(*tasks, return_exceptions=True)
+
+ # Verify final balance is consistent
+ await integration_session.refresh(api_key)
+ # Balance should have some deduction but exact amount depends on implementation
+ assert api_key.balance < 10000
+ assert api_key.balance >= 0 # Should never go negative
+
+
+class TestConcurrentOperations:
+ """Test database consistency under concurrent operations"""
+
+ @pytest.mark.asyncio
+ async def test_multiple_requests_same_api_key(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test multiple concurrent requests with the same API key"""
+ # Mock the wallet info endpoint to track concurrent calls
+ call_count = 0
+ call_times = []
+
+ async def track_concurrent_calls() -> Dict[str, int]:
+ nonlocal call_count
+ call_count += 1
+ call_times.append(time.time())
+ await asyncio.sleep(0.1) # Simulate processing time
+ return {"balance": 1000}
+
+ # Make 10 concurrent requests
+ tasks = []
+ for _ in range(10):
+ task = authenticated_client.get("/v1/wallet/info")
+ tasks.append(task)
+
+ responses = await asyncio.gather(*tasks)
+
+ # All requests should succeed
+ for response in responses:
+ assert response.status_code == 200
+
+ @pytest.mark.asyncio
+ async def test_simultaneous_topup_and_usage(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test simultaneous top-up and balance usage operations"""
+ # Get API key info
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ # API key format is "sk-{hashed_key}"
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ # Set initial balance
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ api_key = result.scalar_one()
+ initial_balance = 5000
+ api_key.balance = initial_balance
+ await integration_session.commit()
+
+ # Mock wallet for topup
+ with patch("router.cashu.wallet") as mock_wallet_func:
+ mock_proof = MagicMock()
+ mock_proof.amount = 2000
+ mock_wallet = AsyncMock()
+ mock_wallet.deserialize_token = AsyncMock(return_value=[mock_proof])
+ mock_wallet.redeem = AsyncMock(return_value=[mock_proof])
+ mock_wallet_func.return_value = mock_wallet
+
+ # Mock proxy endpoint to simulate usage
+ with patch("httpx.AsyncClient.request") as mock_request:
+ # Mock successful proxy response
+ mock_response = AsyncMock()
+ mock_response.status_code = 200
+ mock_response.headers = {"content-type": "application/json"}
+ mock_response.aiter_bytes = AsyncMock(
+ return_value=iter([b'{"result": "ok"}'])
+ )
+ mock_response.is_stream_consumed = False
+ mock_request.return_value = mock_response
+
+ # Run topup and usage concurrently
+ async def topup() -> Any:
+ return await authenticated_client.post(
+ "/v1/wallet/topup", params={"cashu_token": "cashuAey..."}
+ )
+
+ async def use_balance() -> Any:
+ # This would normally deduct balance
+ return await authenticated_client.post(
+ "/v1/chat/completions", json={"model": "test", "messages": []}
+ )
+
+ # Execute concurrently
+ results = await asyncio.gather(
+ topup(), use_balance(), return_exceptions=True
+ )
+ topup_result = results[0]
+ usage_result = results[1]
+
+ # At least one should succeed
+ assert not isinstance(topup_result, Exception) or not isinstance(
+ usage_result, Exception
+ )
+
+ # Verify final balance is consistent
+ await integration_session.refresh(api_key)
+ # Balance should be between initial and initial + topup amount
+ assert api_key.balance >= initial_balance
+ assert api_key.balance <= initial_balance + 2000
+
+ @pytest.mark.asyncio
+ async def test_race_condition_prevention(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test that race conditions are prevented in balance updates"""
+ # Get API key info
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ # API key format is "sk-{hashed_key}"
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ # Set a specific balance
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ api_key = result.scalar_one()
+ api_key.balance = 1000
+ api_key.total_spent = 0
+ api_key.total_requests = 0
+ await integration_session.commit()
+
+ # Create a controlled race condition scenario
+ balance_checks: List[int] = []
+
+ async def check_and_update_balance() -> bool:
+ # Read current balance
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ current_api_key = result.scalar_one()
+ current_balance = current_api_key.balance
+ balance_checks.append(current_balance)
+
+ # Simulate processing delay
+ await asyncio.sleep(0.01)
+
+ # Try to update based on read value
+ current_api_key.balance = current_balance - 100
+ current_api_key.total_spent += 100
+ current_api_key.total_requests += 1
+
+ try:
+ await integration_session.commit()
+ return True
+ except Exception:
+ await integration_session.rollback()
+ return False
+
+ # Run multiple concurrent updates
+ tasks = [check_and_update_balance() for _ in range(5)]
+ results = await asyncio.gather(*tasks, return_exceptions=True)
+
+ # Refresh and check final state
+ await integration_session.refresh(api_key)
+
+ # At least some updates should succeed
+ successful_updates = sum(1 for r in results if r is True)
+ assert successful_updates > 0
+
+ # Final balance should reflect successful updates
+ expected_balance = 1000 - (successful_updates * 100)
+ assert api_key.balance == expected_balance
+ assert api_key.total_spent == successful_updates * 100
+ assert api_key.total_requests == successful_updates
+
+
+class TestDataIntegrity:
+ """Test data integrity constraints and validations"""
+
+ @pytest.mark.asyncio
+ async def test_balance_never_negative(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test that balance can never go negative"""
+ # Get API key info
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ # API key format is "sk-{hashed_key}"
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ # Set low balance
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ api_key = result.scalar_one()
+ api_key.balance = 100
+ await integration_session.commit()
+
+ # Try to refund more than balance
+ response = await authenticated_client.post(
+ "/v1/wallet/refund", json={"amount": 1000}
+ )
+
+ # Should fail
+ assert response.status_code == 400
+ assert "Balance too small to refund" in response.json()["detail"]
+
+ # Verify balance unchanged
+ await integration_session.refresh(api_key)
+ assert api_key.balance == 100
+
+ @pytest.mark.asyncio
+ async def test_primary_key_uniqueness(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test that primary key constraints are enforced"""
+ # Get existing API key hash from authenticated client
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ # Try to manually insert duplicate key with same hash
+ duplicate_key = ApiKey(
+ hashed_key=api_key_hash, balance=5000, total_spent=0, total_requests=0
+ )
+
+ integration_session.add(duplicate_key)
+
+ # Should raise integrity error
+ with pytest.raises(IntegrityError):
+ await integration_session.commit()
+
+ await integration_session.rollback()
+
+ @pytest.mark.asyncio
+ async def test_timestamp_consistency(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test that timestamps are consistent and properly ordered"""
+ # Track request times
+ request_times: List[float] = []
+
+ # Make several requests with delays
+ for i in range(3):
+ start_time = time.time()
+ response = await authenticated_client.get("/v1/wallet/info")
+ assert response.status_code == 200
+ request_times.append(start_time)
+ await asyncio.sleep(0.1)
+
+ # Verify timestamps are monotonically increasing
+ for i in range(1, len(request_times)):
+ assert request_times[i] > request_times[i - 1]
+
+ @pytest.mark.asyncio
+ async def test_numeric_field_constraints(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test constraints on numeric fields"""
+ # Get API key
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ # API key format is "sk-{hashed_key}"
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ api_key = result.scalar_one()
+
+ # Test setting invalid values directly
+ # These should maintain integrity
+ assert api_key.balance >= 0
+ assert api_key.total_spent >= 0
+ assert api_key.total_requests >= 0
+
+ # Verify calculations are consistent
+ if api_key.total_requests > 0:
+ average_cost = api_key.total_spent / api_key.total_requests
+ assert average_cost >= 0
+
+
+class TestPerformance:
+ """Test database performance characteristics"""
+
+ @pytest.mark.asyncio
+ async def test_operation_latency(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test that database operations complete within acceptable time"""
+ operation_times: Dict[str, List[float]] = {
+ "select": [],
+ "update": [],
+ "insert": [],
+ }
+
+ # Test SELECT performance
+ for _ in range(10):
+ start = time.time()
+ response = await authenticated_client.get("/v1/wallet/info")
+ end = time.time()
+ assert response.status_code == 200
+ operation_times["select"].append((end - start) * 1000) # Convert to ms
+
+ # Test UPDATE performance (via topup)
+ with patch("router.cashu.wallet") as mock_wallet_func:
+ mock_proof = MagicMock()
+ mock_proof.amount = 100
+ mock_wallet = AsyncMock()
+ mock_wallet.deserialize_token = AsyncMock(return_value=[mock_proof])
+ mock_wallet.redeem = AsyncMock(return_value=[mock_proof])
+ mock_wallet_func.return_value = mock_wallet
+
+ for _ in range(5):
+ start = time.time()
+ response = await authenticated_client.post(
+ "/v1/wallet/topup", params={"cashu_token": "cashuAey..."}
+ )
+ end = time.time()
+ # Skip if token is invalid (400)
+ if response.status_code == 400:
+ continue
+ assert response.status_code == 200
+ operation_times["update"].append((end - start) * 1000)
+
+ # Verify all operations < 100ms
+ for op_type, times in operation_times.items():
+ if times: # Only check if we have measurements
+ avg_time = sum(times) / len(times)
+ max_time = max(times)
+
+ # Average should be well under 100ms
+ assert avg_time < 100, (
+ f"{op_type} average time {avg_time}ms exceeds 100ms"
+ )
+
+ # No single operation should exceed 200ms
+ assert max_time < 200, f"{op_type} max time {max_time}ms exceeds 200ms"
+
+ @pytest.mark.asyncio
+ async def test_connection_pool_behavior(
+ self,
+ authenticated_client: AsyncClient,
+ integration_app: Any,
+ ) -> None:
+ """Test database connection pool behavior under load"""
+
+ # Make many concurrent requests to test connection pooling
+ async def make_request() -> Response:
+ return await authenticated_client.get("/v1/wallet/info")
+
+ # Create 50 concurrent requests
+ tasks = [make_request() for _ in range(50)]
+
+ start = time.time()
+ responses = await asyncio.gather(*tasks, return_exceptions=True)
+ end = time.time()
+
+ # All should succeed
+ success_count = sum(
+ 1
+ for r in responses
+ if not isinstance(r, Exception)
+ and hasattr(r, "status_code")
+ and r.status_code == 200
+ )
+ assert success_count == 50, f"Only {success_count}/50 requests succeeded"
+
+ # Should complete reasonably quickly (< 5 seconds for 50 requests)
+ total_time = end - start
+ assert total_time < 5.0, f"50 concurrent requests took {total_time}s"
+
+ @pytest.mark.asyncio
+ async def test_index_usage(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test that database indexes are used efficiently"""
+ # Get API key for testing
+ api_key_header = authenticated_client.headers["Authorization"].replace(
+ "Bearer ", ""
+ )
+ # API key format is "sk-{hashed_key}"
+ api_key_hash = (
+ api_key_header[3:] if api_key_header.startswith("sk-") else api_key_header
+ )
+
+ # Primary key lookup should be fast
+ start = time.time()
+ stmt = select(ApiKey).where(ApiKey.hashed_key == api_key_hash) # type: ignore[arg-type]
+ result = await integration_session.execute(stmt)
+ api_key = result.scalar_one()
+ end = time.time()
+
+ lookup_time = (end - start) * 1000
+ assert lookup_time < 10, f"Primary key lookup took {lookup_time}ms"
+
+ # Verify we got the right record
+ assert api_key.hashed_key == api_key_hash
diff --git a/tests/integration/test_error_handling_edge_cases.py b/tests/integration/test_error_handling_edge_cases.py
new file mode 100644
index 00000000..d8f48aa9
--- /dev/null
+++ b/tests/integration/test_error_handling_edge_cases.py
@@ -0,0 +1,662 @@
+"""Comprehensive error handling and edge case tests"""
+
+import asyncio
+import time
+from typing import Any
+from unittest.mock import AsyncMock, MagicMock, patch
+
+import pytest
+from httpx import AsyncClient, ConnectError
+from sqlalchemy.ext.asyncio import AsyncSession
+from sqlmodel import select
+
+from router.db import ApiKey
+
+
+class TestNetworkFailureScenarios:
+ """Test various network failure scenarios"""
+
+ @pytest.mark.asyncio
+ async def test_mint_service_unavailable(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test behavior when mint service is unavailable"""
+ # Get the existing mock wallet from the fixture
+ from router.cashu import wallet
+
+ mock_wallet = wallet()
+
+ # Temporarily override the send method to simulate failure
+ with patch.object(
+ mock_wallet,
+ "send",
+ AsyncMock(side_effect=ConnectError("Mint service unavailable")),
+ ):
+ # Try to refund when mint is down
+ response = await authenticated_client.post(
+ "/v1/wallet/refund", json={"amount": 1000}
+ )
+
+ # Should get error response (503 for service unavailable)
+ assert response.status_code == 503
+ # The error detail might vary based on implementation
+
+ @pytest.mark.asyncio
+ async def test_upstream_llm_service_down(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test proxy behavior when upstream LLM service is down"""
+ # Mock at the router level to simulate upstream being down
+ with patch("router.proxy.httpx.AsyncClient") as mock_client_class:
+ # Create a mock client instance
+ mock_client = AsyncMock()
+ mock_client_class.return_value = mock_client
+ mock_client.__aenter__ = AsyncMock(return_value=mock_client)
+ mock_client.__aexit__ = AsyncMock(return_value=None)
+ mock_client.aclose = AsyncMock()
+
+ # Make the send method raise ConnectError
+ mock_client.send = AsyncMock(side_effect=ConnectError("Connection refused"))
+ mock_client.build_request = MagicMock(return_value=MagicMock())
+
+ # Try to make a proxy request
+ response = await authenticated_client.post(
+ "/v1/chat/completions",
+ json={
+ "model": "gpt-3.5-turbo",
+ "messages": [{"role": "user", "content": "Hello"}],
+ },
+ )
+
+ # Should get appropriate error (502 for upstream error)
+ assert response.status_code == 502
+ # Error detail depends on implementation
+
+ @pytest.mark.asyncio
+ async def test_partial_request_failures(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test handling of partial failures during streaming"""
+
+ # Mock streaming response that fails midway
+ async def mock_aiter_bytes() -> Any: # type: ignore[misc]
+ yield b'data: {"choices": [{"delta": {"content": "Hello"}}]}\n\n'
+ yield b'data: {"choices": [{"delta": {"content": " World"}}]}\n\n'
+ raise ConnectError("Connection lost")
+
+ with patch("httpx.AsyncClient.request") as mock_request:
+ 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.is_stream_consumed = False
+ mock_request.return_value = mock_response
+
+ # Make streaming request
+ response = await authenticated_client.post(
+ "/v1/chat/completions",
+ json={
+ "model": "gpt-3.5-turbo",
+ "messages": [{"role": "user", "content": "Hello"}],
+ "stream": True,
+ },
+ )
+
+ # Should still return 200 even with partial failure
+ # The streaming error happens after headers are sent
+ assert response.status_code == 200
+
+ # In real implementation, partial charges would be handled
+ # but our mock doesn't actually deduct balance
+
+ @pytest.mark.asyncio
+ async def test_timeout_handling(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test request timeout handling"""
+ # Similar to above, we test timeout handling exists
+ # but can't easily trigger real timeouts in test environment
+
+ with patch("httpx.AsyncClient.send") as mock_send:
+ # Create a mock timeout response
+ mock_response = AsyncMock()
+ mock_response.status_code = 504
+ mock_response.headers = {"content-type": "application/json"}
+ mock_response.json.return_value = {"error": "Gateway Timeout"}
+ mock_response.text = '{"error": "Gateway Timeout"}'
+ mock_response.content = b'{"error": "Gateway Timeout"}'
+ mock_response.aiter_bytes = AsyncMock(
+ return_value=AsyncMock(
+ __aiter__=lambda self: self,
+ __anext__=AsyncMock(side_effect=StopAsyncIteration),
+ )
+ )
+ mock_send.return_value = mock_response
+
+ # Make request
+ response = await authenticated_client.post(
+ "/v1/chat/completions",
+ json={
+ "model": "gpt-3.5-turbo",
+ "messages": [{"role": "user", "content": "Hello"}],
+ },
+ )
+
+ # Should pass through the error
+ assert response.status_code >= 500
+
+
+class TestInvalidInputHandling:
+ """Test handling of various invalid inputs"""
+
+ @pytest.mark.asyncio
+ async def test_malformed_cashu_tokens(
+ self,
+ authenticated_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test various malformed Cashu token formats"""
+ malformed_tokens = [
+ "", # Empty token
+ "not-a-token", # Invalid format
+ "cashu", # Incomplete
+ "cashuA" + "x" * 10000, # Extremely long
+ "cashuA" + "\x00" + "test", # Null bytes
+ "cashuA" + "\n\r" + "test", # Control characters
+ "cashuAeyJhbGciOi", # Truncated base64
+ "cashuA!!!invalid-base64!!!", # Invalid base64
+ ]
+
+ for token in malformed_tokens:
+ response = await authenticated_client.post(
+ "/v1/wallet/topup", params={"cashu_token": token}
+ )
+
+ # All should fail with 400
+ assert response.status_code == 400, f"Token {repr(token)} should fail"
+ assert "invalid" in response.json()["detail"].lower()
+
+ @pytest.mark.asyncio
+ async def test_invalid_json_payloads(
+ self,
+ authenticated_client: AsyncClient,
+ ) -> None:
+ """Test handling of invalid JSON in requests"""
+ # Test malformed JSON
+ response = await authenticated_client.post(
+ "/v1/chat/completions",
+ content='{"model": "gpt-3.5-turbo", "messages": [}', # Invalid JSON
+ headers={"content-type": "application/json"},
+ )
+ assert response.status_code in [
+ 400,
+ 422,
+ ] # Either is acceptable for malformed JSON
+
+ # Test wrong content type
+ response = await authenticated_client.post(
+ "/v1/chat/completions",
+ content="not json at all",
+ headers={"content-type": "application/json"},
+ )
+ assert response.status_code in [400, 422]
+
+ # Test missing required fields - proxy endpoints just forward, so might get different error
+ response = await authenticated_client.post(
+ "/v1/chat/completions",
+ json={"model": "gpt-3.5-turbo"}, # Missing messages
+ )
+ assert response.status_code >= 400 # Any 4xx error is acceptable
+
+ @pytest.mark.asyncio
+ async def test_sql_injection_attempts(
+ self,
+ integration_client: AsyncClient,
+ integration_session: AsyncSession,
+ ) -> None:
+ """Test that SQL injection attempts are properly handled"""
+ # SQL injection attempts in various places
+ injection_payloads = [
+ "'; DROP TABLE api_keys; --",
+ "1' OR '1'='1",
+ "admin'--",
+ "1; UPDATE api_keys SET balance=999999999;",
+ "' UNION SELECT * FROM api_keys--",
+ ]
+
+ for payload in injection_payloads:
+ # Try injection in authorization header
+ response = await integration_client.get(
+ "/v1/wallet/info", headers={"Authorization": f"Bearer {payload}"}
+ )
+ assert response.status_code == 401
+
+ # Try injection in refund amount
+ response = await integration_client.post(
+ "/v1/wallet/refund", json={"amount": payload}
+ )
+ assert response.status_code in [
+ 401,
+ 422,
+ ] # Unauthorized or validation error
+
+ @pytest.mark.asyncio
+ async def test_xss_in_headers_params(
+ self,
+ authenticated_client: AsyncClient,
+ ) -> None:
+ """Test XSS prevention in headers and parameters"""
+ xss_payloads = [
+ "",
+ "javascript:alert(1)",
+ "
",
+ "