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)", + "", + "", + "'+alert(1)+'", + ] + + for payload in xss_payloads: + # Try XSS in custom headers + response = await authenticated_client.get( + "/v1/wallet/info", headers={"X-Custom-Header": payload} + ) + # Should process normally, but payload should be escaped/ignored + assert response.status_code == 200 + + # If response includes headers, verify they're escaped + if "X-Custom-Header" in response.headers: + assert "" in html_content and "" in html_content + assert "" in html_content and "" in html_content + + # Should have CSS styling + assert "