diff --git a/pyproject.toml b/pyproject.toml index 936a6f5e..42daeeaa 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,4 +14,4 @@ dependencies = [ ] [dependency-groups] -dev = ["mypy>=1.15.0", "ruff>=0.11.6", "openai>=1.76.0"] +dev = ["mypy>=1.15.0", "ruff>=0.11.6", "openai>=1.76.0", "pytest>=8.0.0", "pytest-asyncio>=0.24.0", "httpx>=0.25.2"] diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 00000000..094801d8 --- /dev/null +++ b/pytest.ini @@ -0,0 +1,17 @@ +[tool:pytest] +testpaths = tests +python_files = test_*.py +python_classes = Test* +python_functions = test_* +asyncio_mode = auto +asyncio_default_fixture_loop_scope = function +addopts = + -v + --tb=short + --strict-markers + --disable-warnings + -p no:warnings +markers = + asyncio: marks tests as async (deselect with '-m "not asyncio"') + integration: marks tests as integration tests + unit: marks tests as unit tests \ No newline at end of file diff --git a/tests/README.md b/tests/README.md new file mode 100644 index 00000000..4ab5abc3 --- /dev/null +++ b/tests/README.md @@ -0,0 +1,63 @@ +# FastAPI Async Unit Tests + +This directory contains async unit tests for the Routstr proxy FastAPI application. + +## Installation + +First, ensure you have the development dependencies installed: + +```bash +uv pip install -e ".[dev]" +``` + +## Running Tests + +To run all tests: +```bash +pytest +``` + +To run tests with coverage: +```bash +pytest --cov=router --cov-report=html +``` + +To run specific test files: +```bash +pytest tests/test_main.py +pytest tests/test_account.py +pytest tests/test_proxy.py +pytest tests/test_models.py +``` + +To run only async tests: +```bash +pytest -m asyncio +``` + +## Test Structure + +- `conftest.py` - Pytest fixtures and configuration +- `test_main.py` - Tests for main app endpoints +- `test_account.py` - Tests for wallet/account management endpoints +- `test_proxy.py` - Tests for the proxy functionality with mocked upstream +- `test_models.py` - Tests for model pricing and data structures + +## Key Fixtures + +- `async_client` - Async HTTP client for testing FastAPI endpoints +- `test_session` - In-memory SQLite database session for tests +- `test_api_key` - Pre-configured API key with balance +- `api_key_with_balance` - API key with sufficient balance for proxy tests + +## Environment Variables + +The tests automatically set up required environment variables in `conftest.py`. No manual configuration needed. + +## Writing New Tests + +1. Use `@pytest.mark.asyncio` for async tests +2. Use the provided fixtures for database and client access +3. Mock external dependencies (like upstream API calls) +4. Test both success and error cases +5. Verify database state changes when applicable \ No newline at end of file diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 00000000..0519ecba --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 00000000..ed98bc35 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,181 @@ +import asyncio +import os +import pytest +import pytest_asyncio +from typing import AsyncGenerator +from fastapi.testclient import TestClient +from httpx import AsyncClient, ASGITransport +from sqlmodel import SQLModel +from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession +from sqlalchemy.orm import sessionmaker +from unittest.mock import patch, MagicMock, AsyncMock + +# Save original environment variables +ORIGINAL_ENV = os.environ.copy() + +# Set test environment variables before importing the app +TEST_ENV = { + "UPSTREAM_BASE_URL": "https://api.example.com", + "UPSTREAM_API_KEY": "test-upstream-key", + "NAME": "TestRoutstrNode", + "DESCRIPTION": "Test Node", + "NPUB": "npub1test", + "MINT": "https://test.mint.com", + "HTTP_URL": "http://test.example.com", + "ONION_URL": "http://test.onion", + "CORS_ORIGINS": "*", + "RECEIVE_LN_ADDRESS": "test@lightning.address", + "COST_PER_REQUEST": "1", + "COST_PER_1K_INPUT_TOKENS": "0", + "COST_PER_1K_OUTPUT_TOKENS": "0", + "MODEL_BASED_PRICING": "false" +} + +# Apply test environment +os.environ.update(TEST_ENV) + +# Mock the cashu wallet initialization before importing +with patch("router.cashu._initialize_wallet") as mock_init_wallet: + mock_wallet = AsyncMock() + mock_wallet.available_balance = 1000 + mock_wallet.proofs = [] + mock_wallet.split = AsyncMock(return_value=([], [])) + mock_init_wallet.return_value = mock_wallet + + with patch("router.cashu.WALLET", mock_wallet): + from router.main import app + from router.db import get_session + from router.models import MODELS + + +@pytest.fixture(scope="session") +def event_loop(): + """Create an instance of the default event loop for the test session.""" + loop = asyncio.get_event_loop_policy().new_event_loop() + yield loop + loop.close() + + +@pytest_asyncio.fixture(scope="function") +async def test_engine(): + """Create a test database engine - new for each test.""" + engine = create_async_engine( + "sqlite+aiosqlite:///:memory:", + echo=False, + future=True, + ) + + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + yield engine + + await engine.dispose() + + +@pytest_asyncio.fixture +async def test_session(test_engine) -> AsyncSession: + """Create a test database session.""" + async_session = sessionmaker( + test_engine, class_=AsyncSession, expire_on_commit=False + ) + + async with async_session() as session: + yield session + + +@pytest.fixture +def test_client() -> TestClient: + """Create a test client for the FastAPI app.""" + with patch.dict(os.environ, TEST_ENV, clear=True): + with patch("router.cashu._initialize_wallet") as mock_init: + mock_wallet = AsyncMock() + mock_wallet.available_balance = 1000 + mock_wallet.proofs = [] + mock_wallet.split = AsyncMock(return_value=([], [])) + mock_init.return_value = mock_wallet + + with patch("router.models.update_sats_pricing") as mock_update: + mock_update.return_value = None + return TestClient(app) + + +@pytest_asyncio.fixture +async def async_client(test_session) -> AsyncClient: + """Create an async test client with dependency overrides.""" + async def override_get_session(): + yield test_session + + app.dependency_overrides[get_session] = override_get_session + + # Mock startup tasks + with patch.dict(os.environ, TEST_ENV, clear=True): + with patch("router.cashu._initialize_wallet") as mock_init: + mock_wallet = AsyncMock() + mock_wallet.available_balance = 1000 + mock_wallet.proofs = [] + mock_wallet.split = AsyncMock(return_value=([], [])) + mock_init.return_value = mock_wallet + + with patch("router.models.update_sats_pricing") as mock_update: + mock_update.return_value = None + + async with AsyncClient( + transport=ASGITransport(app=app), + base_url="http://test" + ) as client: + yield client + + app.dependency_overrides.clear() + + +@pytest.fixture +def mock_models(): + """Mock models data for testing.""" + return [ + { + "id": "gpt-4", + "name": "GPT-4", + "created": 1680000000, + "description": "Test model", + "context_length": 8192, + "architecture": { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "cl100k_base", + "instruct_type": "none" + }, + "pricing": { + "prompt": 0.03, + "completion": 0.06, + "request": 0.001, + "image": 0.0, + "web_search": 0.0, + "internal_reasoning": 0.0 + }, + "top_provider": { + "context_length": 8192, + "max_completion_tokens": 4096, + "is_moderated": False + } + } + ] + + +# Cleanup after all tests +@pytest.fixture(scope="session", autouse=True) +def cleanup(): + yield + # Restore original environment carefully + current_keys = set(os.environ.keys()) + original_keys = set(ORIGINAL_ENV.keys()) + + # Remove keys that weren't in original + for key in current_keys - original_keys: + if key != 'PYTEST_CURRENT_TEST': # Don't touch pytest's own variables + os.environ.pop(key, None) + + # Restore original values + for key, value in ORIGINAL_ENV.items(): + os.environ[key] = value \ No newline at end of file diff --git a/tests/test_account.py b/tests/test_account.py new file mode 100644 index 00000000..bdcb8f94 --- /dev/null +++ b/tests/test_account.py @@ -0,0 +1,217 @@ +import pytest +import pytest_asyncio +import hashlib +import uuid +from unittest.mock import patch, AsyncMock, MagicMock +from httpx import AsyncClient +from router.db import ApiKey, AsyncSession + + +def hash_api_key(api_key: str) -> str: + """Hash an API key for storage.""" + return hashlib.sha256(api_key.encode()).hexdigest() + + +@pytest_asyncio.fixture +async def test_api_key(test_session: AsyncSession) -> ApiKey: + """Create a test API key in the database.""" + # Use unique key for each test + unique_id = str(uuid.uuid4())[:8] + api_key = f"test-api-key-{unique_id}" + + key = ApiKey( + hashed_key=api_key, + balance=1000000, # 1000 sats in msats + refund_address="test@lightning.address", + total_spent=0, + total_requests=0 + ) + + test_session.add(key) + await test_session.commit() + await test_session.refresh(key) + + return key + + +@pytest.mark.asyncio +async def test_account_info_with_valid_key( + async_client: AsyncClient, + test_api_key: ApiKey +): + """Test getting account info with a valid API key.""" + response = await async_client.get( + "/v1/wallet/", + headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"} + ) + + assert response.status_code == 200 + data = response.json() + + assert data["api_key"] == f"sk-{test_api_key.hashed_key}" + assert data["balance"] == 1000000 + + +@pytest.mark.asyncio +async def test_account_info_without_auth(async_client: AsyncClient): + """Test that account info requires authentication.""" + response = await async_client.get("/v1/wallet/") + + assert response.status_code == 422 # Missing required header + + +@pytest.mark.asyncio +async def test_account_info_with_invalid_key(async_client: AsyncClient): + """Test account info with an invalid API key.""" + response = await async_client.get( + "/v1/wallet/", + headers={"Authorization": "Bearer invalid-key"} + ) + + assert response.status_code == 401 + + +@pytest.mark.asyncio +async def test_refund_balance_with_address( + async_client: AsyncClient, + test_api_key: ApiKey, + test_session: AsyncSession +): + """Test refunding balance when refund address is set.""" + # Need to patch the refund_balance at the module level to intercept the call + with patch("router.account.refund_balance", new_callable=AsyncMock) as mock_refund: + mock_refund.return_value = 1000000 + + response = await async_client.post( + "/v1/wallet/refund", + headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"} + ) + + assert response.status_code == 200 + data = response.json() + + assert data["recipient"] == "test@lightning.address" + assert data["msats"] == 1000000 + + # Verify balance was zeroed + await test_session.refresh(test_api_key) + assert test_api_key.balance == 0 + + # Verify refund_balance was called + mock_refund.assert_called_once() + + +@pytest.mark.asyncio +async def test_refund_balance_without_address( + async_client: AsyncClient, + test_session: AsyncSession +): + """Test refunding balance when no refund address is set.""" + # Create key without refund address - with unique ID + unique_id = str(uuid.uuid4())[:8] + api_key = f"test-key-no-refund-{unique_id}" + + key = ApiKey( + hashed_key=api_key, + balance=500000, + refund_address=None, + total_spent=0, + total_requests=0 + ) + + test_session.add(key) + await test_session.commit() + + # Mock at the router.account module level + with patch("router.account.create_token", new_callable=AsyncMock) as mock_create_token: + mock_create_token.return_value = "cashuBqQSEQ..." + + response = await async_client.post( + "/v1/wallet/refund", + headers={"Authorization": f"Bearer sk-{api_key}"} + ) + + assert response.status_code == 200 + data = response.json() + + assert data["recipient"] is None + assert data["msats"] == 500000 + assert data["token"] == "cashuBqQSEQ..." + + # Verify create_token was called with the correct amount + mock_create_token.assert_called_once_with(500000) + + +@pytest.mark.asyncio +async def test_topup_balance_endpoint( + async_client: AsyncClient, + test_api_key: ApiKey, + test_session: AsyncSession +): + """Test topping up balance with a cashu token.""" + # Mock at the router.account module level to intercept the import + with patch("router.account.credit_balance", new_callable=AsyncMock) as mock_credit: + mock_credit.return_value = {"msats": 500000} + + response = await async_client.post( + "/v1/wallet/topup?cashu_token=cashuBqQSEQ...", + headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"} + ) + + assert response.status_code == 200 + data = response.json() + assert data == {"msats": 500000} + + # Verify credit_balance was called + mock_credit.assert_called_once() + + +@pytest.mark.asyncio +async def test_topup_balance_requires_cashu_token( + async_client: AsyncClient, + test_api_key: ApiKey +): + """Test that topup endpoint requires a cashu token.""" + response = await async_client.post( + "/v1/wallet/topup", + headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"}, + json={} + ) + + assert response.status_code == 422 # Missing required field + + +@pytest.mark.asyncio +async def test_account_with_cashu_token( + async_client: AsyncClient, + test_session: AsyncSession +): + """Test authentication with a cashu token creates a new account.""" + cashu_token = "cashuBqQSEQ123456" + hashed = hash_api_key(cashu_token) + + with patch("router.cashu.credit_balance", new_callable=AsyncMock) as mock_credit: + # Mock successful token redemption + mock_credit.return_value = 5000000 # 5000 sats + + # Mock token deserialization + with patch("router.cashu.deserialize_token_from_string") as mock_deserialize: + mock_token = MagicMock() + mock_token.mint = "https://test.mint.com" + mock_deserialize.return_value = mock_token + + # Mock wallet receive + with patch("router.cashu._handle_token_receive", new_callable=AsyncMock) as mock_receive: + mock_receive.return_value = 5000000 + + response = await async_client.get( + "/v1/wallet/", + headers={"Authorization": f"Bearer {cashu_token}"} + ) + + assert response.status_code == 200 + data = response.json() + + # Check that a new key was created with the hashed token + assert data["api_key"].startswith("sk-") + assert data["balance"] >= 0 # Balance should be set after credit_balance \ No newline at end of file diff --git a/tests/test_main.py b/tests/test_main.py new file mode 100644 index 00000000..0a66632a --- /dev/null +++ b/tests/test_main.py @@ -0,0 +1,63 @@ +import pytest +from httpx import AsyncClient +from unittest.mock import patch + + +@pytest.mark.asyncio +async def test_root_endpoint(async_client: AsyncClient): + """Test the root endpoint returns expected information.""" + # Mock the environment variables for this specific test + with patch("os.environ.get") as mock_env_get: + def env_side_effect(key, default=None): + env_map = { + "NAME": "TestRoutstrNode", + "DESCRIPTION": "Test Node", + "NPUB": "npub1test", + "MINT": "https://test.mint.com", + "HTTP_URL": "http://test.example.com", + "ONION_URL": "http://test.onion", + } + return env_map.get(key, default) + + mock_env_get.side_effect = env_side_effect + + response = await async_client.get("/") + + assert response.status_code == 200 + data = response.json() + + # The app reads from env vars during import, so check what we actually get + assert "name" in data + assert "description" in data + assert data["version"] == "0.0.1" + assert "npub" in data + assert "mint" in data + assert "http_url" in data + assert "onion_url" in data + assert "models" in data + + +@pytest.mark.asyncio +async def test_cors_headers(async_client: AsyncClient): + """Test that CORS headers are properly set.""" + response = await async_client.options( + "/", + headers={ + "Origin": "http://localhost:3000", + "Access-Control-Request-Method": "GET", + } + ) + + assert response.status_code == 200 + # Check that CORS is working (might be * or specific origin) + assert "access-control-allow-origin" in response.headers + assert "GET" in response.headers["access-control-allow-methods"] + + +@pytest.mark.asyncio +async def test_startup_event_initializes_properly(test_client): + """Test that the startup event runs without errors.""" + # The test_client fixture already triggers the startup event + # This test ensures no exceptions are raised during startup + response = test_client.get("/") + assert response.status_code == 200 \ No newline at end of file diff --git a/tests/test_models.py b/tests/test_models.py new file mode 100644 index 00000000..d28a2f2f --- /dev/null +++ b/tests/test_models.py @@ -0,0 +1,200 @@ +import pytest +import asyncio +from unittest.mock import patch, AsyncMock, MagicMock +from router.models import Model, Architecture, Pricing, TopProvider, update_sats_pricing, MODELS + + +@pytest.fixture +def sample_model() -> Model: + """Create a sample model for testing.""" + return Model( + id="test-model", + name="Test Model", + created=1700000000, + description="A test model", + context_length=4096, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="test_tokenizer", + instruct_type="chat" + ), + pricing=Pricing( + prompt=0.01, + completion=0.02, + request=0.001, + image=0.0, + web_search=0.0, + internal_reasoning=0.0 + ), + top_provider=TopProvider( + context_length=4096, + max_completion_tokens=2048, + is_moderated=False + ) + ) + + +@pytest.mark.asyncio +async def test_update_sats_pricing_calculation(sample_model: Model): + """Test that sats pricing is calculated correctly.""" + # Mock the sats_usd_ask_price function + with patch("router.models.sats_usd_ask_price", new_callable=AsyncMock) as mock_price: + mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD + + # Temporarily replace MODELS + original_models = MODELS[:] + MODELS.clear() + MODELS.append(sample_model) + + # Run one iteration of the pricing update + sleep_called = asyncio.Event() + + async def mock_sleep(duration): + sleep_called.set() + raise asyncio.CancelledError() + + with patch("asyncio.sleep", side_effect=mock_sleep): + try: + # Create and run the task + task = asyncio.create_task(update_sats_pricing()) + + # Wait for the first iteration to complete + await sleep_called.wait() + + # Check that sats pricing was calculated + assert sample_model.sats_pricing is not None + + # Verify calculations (prices in USD / sats_to_usd) + assert sample_model.sats_pricing.prompt == 0.01 / 0.0001 # 100 sats + assert sample_model.sats_pricing.completion == 0.02 / 0.0001 # 200 sats + assert sample_model.sats_pricing.request == 0.001 / 0.0001 # 10 sats + + # Verify max_cost calculation for model with top_provider + expected_max_context = 4096 * sample_model.sats_pricing.prompt + expected_max_completion = 2048 * sample_model.sats_pricing.completion + assert sample_model.sats_pricing.max_cost == expected_max_context + expected_max_completion + + except asyncio.CancelledError: + pass + finally: + # Restore original models + MODELS.clear() + MODELS.extend(original_models) + + +@pytest.mark.asyncio +async def test_update_sats_pricing_without_top_provider(): + """Test sats pricing calculation for models without top_provider.""" + model_without_top = Model( + id="test-model-no-top", + name="Test Model No Top", + created=1700000000, + description="A test model without top provider", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="test_tokenizer", + instruct_type=None + ), + pricing=Pricing( + prompt=0.01, + completion=0.02, + request=0.001, + image=0.01, + web_search=0.005, + internal_reasoning=0.015 + ), + top_provider=None + ) + + with patch("router.models.sats_usd_ask_price", new_callable=AsyncMock) as mock_price: + mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD + + original_models = MODELS[:] + MODELS.clear() + MODELS.append(model_without_top) + + sleep_called = asyncio.Event() + + async def mock_sleep(duration): + sleep_called.set() + raise asyncio.CancelledError() + + with patch("asyncio.sleep", side_effect=mock_sleep): + try: + task = asyncio.create_task(update_sats_pricing()) + await sleep_called.wait() + + assert model_without_top.sats_pricing is not None + + # Verify the fallback max_cost calculation + p = model_without_top.sats_pricing.prompt * 1_000_000 + c = model_without_top.sats_pricing.completion * 32_000 + r = model_without_top.sats_pricing.request * 100_000 + i = model_without_top.sats_pricing.image * 100 + w = model_without_top.sats_pricing.web_search * 1000 + ir = model_without_top.sats_pricing.internal_reasoning * 100 + expected_max = p + c + r + i + w + ir + + assert model_without_top.sats_pricing.max_cost == expected_max + + except asyncio.CancelledError: + pass + finally: + MODELS.clear() + MODELS.extend(original_models) + + +@pytest.mark.asyncio +async def test_update_sats_pricing_handles_errors(): + """Test that update_sats_pricing handles errors gracefully.""" + with patch("router.models.sats_usd_ask_price", new_callable=AsyncMock) as mock_price: + mock_price.side_effect = Exception("API Error") + + error_printed = False + original_print = print + + def mock_print(msg): + nonlocal error_printed + if isinstance(msg, Exception) and str(msg) == "API Error": + error_printed = True + original_print(msg) + + with patch("builtins.print", side_effect=mock_print): + sleep_called = asyncio.Event() + + async def mock_sleep(duration): + sleep_called.set() + raise asyncio.CancelledError() + + with patch("asyncio.sleep", side_effect=mock_sleep): + try: + task = asyncio.create_task(update_sats_pricing()) + await sleep_called.wait() + + # Verify error was printed + assert error_printed + + except asyncio.CancelledError: + pass + + +def test_model_serialization(sample_model: Model): + """Test that models can be serialized and deserialized correctly.""" + model_dict = sample_model.dict() + + # Verify all fields are present + assert model_dict["id"] == "test-model" + assert model_dict["name"] == "Test Model" + assert model_dict["pricing"]["prompt"] == 0.01 + assert model_dict["architecture"]["modality"] == "text" + assert model_dict["top_provider"]["context_length"] == 4096 + + # Test deserialization + new_model = Model(**model_dict) + assert new_model.id == sample_model.id + assert new_model.pricing.prompt == sample_model.pricing.prompt \ No newline at end of file diff --git a/tests/test_proxy.py b/tests/test_proxy.py new file mode 100644 index 00000000..603eb522 --- /dev/null +++ b/tests/test_proxy.py @@ -0,0 +1,334 @@ +import pytest +import pytest_asyncio +import json +import os +import uuid +from unittest.mock import AsyncMock, patch +from httpx import AsyncClient, Response as HttpxResponse +from router.db import ApiKey, AsyncSession + + +@pytest_asyncio.fixture +async def api_key_with_balance(test_session: AsyncSession) -> ApiKey: + """Create an API key with sufficient balance.""" + unique_id = str(uuid.uuid4())[:8] + key = ApiKey( + hashed_key=f"test-hashed-key-{unique_id}", + balance=10000000, # 10,000 sats in msats + refund_address=None, + total_spent=0, + total_requests=0 + ) + test_session.add(key) + await test_session.commit() + await test_session.refresh(key) + return key + + +@pytest.mark.asyncio +async def test_proxy_requires_authentication(async_client: AsyncClient): + """Test that proxy endpoints require authentication.""" + response = await async_client.post("/v1/chat/completions") + + assert response.status_code == 401 + assert "API key or Cashu token required" in response.json()["detail"]["error"]["message"] + + +@pytest.mark.asyncio +async def test_proxy_with_insufficient_balance( + async_client: AsyncClient, + test_session: AsyncSession +): + """Test proxy request with insufficient balance.""" + # Create key with minimal balance + unique_id = str(uuid.uuid4())[:8] + key = ApiKey( + hashed_key=f"low-balance-key-{unique_id}", + balance=100, # Only 0.1 sats + refund_address=None, + total_spent=0, + total_requests=0 + ) + test_session.add(key) + await test_session.commit() + + # Mock the models.json check + with patch("os.path.exists", return_value=False): + response = await async_client.post( + "/v1/chat/completions", + headers={"Authorization": f"Bearer sk-{key.hashed_key}"}, + json={"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]} + ) + + assert response.status_code == 402 + assert "Insufficient balance" in response.json()["detail"]["error"]["message"] + + +@pytest.mark.asyncio +async def test_proxy_invalid_json_body( + async_client: AsyncClient, + api_key_with_balance: ApiKey +): + """Test proxy request with invalid JSON body.""" + response = await async_client.post( + "/v1/chat/completions", + headers={ + "Authorization": f"Bearer sk-{api_key_with_balance.hashed_key}", + "Content-Type": "application/json" + }, + content=b'{"invalid": json",}' # Invalid JSON + ) + + assert response.status_code == 400 + error_data = response.json() + assert error_data["error"]["type"] == "invalid_request_error" + assert error_data["error"]["code"] == "invalid_json" + + +@pytest.mark.asyncio +async def test_proxy_successful_request_mock( + async_client: AsyncClient, + api_key_with_balance: ApiKey, + test_session: AsyncSession +): + """Test successful proxy request with mocked upstream.""" + mock_response_data = { + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "gpt-4", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello! How can I help you?"}, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 9, + "completion_tokens": 10, + "total_tokens": 19 + } + } + + with patch("httpx.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + + # Create a mock response + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.aread = AsyncMock(return_value=json.dumps(mock_response_data).encode()) + mock_response.aiter_bytes = AsyncMock() + mock_response.aclose = AsyncMock() + + mock_client.send = AsyncMock(return_value=mock_response) + mock_client.build_request = AsyncMock() + mock_client.aclose = AsyncMock() + + # Also mock the models.json check and pay_out + with patch("os.path.exists", return_value=False): + with patch("router.cashu.pay_out_with_new_session") as mock_payout: + mock_payout.return_value = None + + response = await async_client.post( + "/v1/chat/completions", + headers={"Authorization": f"Bearer sk-{api_key_with_balance.hashed_key}"}, + json={ + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}] + } + ) + + assert response.status_code == 200 + response_json = response.json() + + # Verify the response includes the original data plus cost + assert response_json["id"] == "chatcmpl-123" + assert "cost" in response_json + assert response_json["cost"]["total_msats"] >= 0 + + # Verify balance was deducted + await test_session.refresh(api_key_with_balance) + assert api_key_with_balance.balance < 10000000 + assert api_key_with_balance.total_requests == 1 + + +@pytest.mark.asyncio +async def test_proxy_streaming_response( + async_client: AsyncClient, + api_key_with_balance: ApiKey +): + """Test proxy request with streaming response.""" + # Mock SSE stream chunks + stream_chunks = [ + b'data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"gpt-4","choices":[{"delta":{"content":"Hello"},"index":0}]}\n\n', + b'data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"gpt-4","choices":[{"delta":{"content":" there!"},"index":0}]}\n\n', + b'data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1677652288,"model":"gpt-4","choices":[{"delta":{},"index":0,"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":3,"total_tokens":12}}\n\n', + b'data: [DONE]\n\n' + ] + + async def mock_aiter_bytes(): + for chunk in stream_chunks: + yield chunk + + with patch("httpx.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + + 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.aclose = AsyncMock() + + mock_client.send = AsyncMock(return_value=mock_response) + mock_client.build_request = AsyncMock() + mock_client.aclose = AsyncMock() + + with patch("os.path.exists", return_value=False): + with patch("router.cashu.pay_out_with_new_session") as mock_payout: + mock_payout.return_value = None + + response = await async_client.post( + "/v1/chat/completions", + headers={"Authorization": f"Bearer sk-{api_key_with_balance.hashed_key}"}, + json={ + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "stream": True + } + ) + + assert response.status_code == 200 + assert response.headers["content-type"] == "text/event-stream" + + +@pytest.mark.asyncio +async def test_proxy_handles_upstream_errors( + async_client: AsyncClient, + api_key_with_balance: ApiKey +): + """Test proxy handles upstream connection errors gracefully.""" + with patch("httpx.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + + # Simulate connection error + mock_client.send.side_effect = Exception("Connection refused") + mock_client.build_request = AsyncMock() + mock_client.aclose = AsyncMock() + + with patch("os.path.exists", return_value=False): + response = await async_client.post( + "/v1/chat/completions", + headers={"Authorization": f"Bearer sk-{api_key_with_balance.hashed_key}"}, + json={ + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}] + } + ) + + assert response.status_code == 500 + error_data = response.json() + assert error_data["error"]["type"] == "internal_error" + assert error_data["error"]["message"] == "An unexpected server error occurred" + + +@pytest.mark.asyncio +async def test_proxy_with_model_based_pricing( + async_client: AsyncClient, + test_session: AsyncSession +): + """Test proxy with model-based pricing enabled.""" + # Create API key with sufficient balance + unique_id = str(uuid.uuid4())[:8] + key = ApiKey( + hashed_key=f"model-pricing-key-{unique_id}", + balance=10000000, # 10,000 sats + refund_address=None, + total_spent=0, + total_requests=0 + ) + test_session.add(key) + await test_session.commit() + + with patch.dict(os.environ, {"MODEL_BASED_PRICING": "true"}): + with patch("os.path.exists", return_value=True): + # Mock a model with pricing + from router.models import MODELS, Model, Pricing, Architecture, TopProvider + + test_model = Model( + id="gpt-4", + name="GPT-4", + created=1680000000, + description="Test model", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="cl100k_base", + instruct_type="none" + ), + pricing=Pricing( + prompt=0.03, + completion=0.06, + request=0.001, + image=0.0, + web_search=0.0, + internal_reasoning=0.0 + ), + sats_pricing=Pricing( + prompt=300, # 300 sats per 1k tokens + completion=600, + request=10, + image=0.0, + web_search=0.0, + internal_reasoning=0.0, + max_cost=5000 # 5000 sats max + ), + top_provider=TopProvider( + context_length=8192, + max_completion_tokens=4096, + is_moderated=False + ) + ) + + # Temporarily replace models + original_models = MODELS[:] + MODELS.clear() + MODELS.append(test_model) + + # Mock the upstream HTTP client + with patch("httpx.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + + # Create a mock response + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.aread = AsyncMock(return_value=b'{"id": "test", "model": "gpt-4"}') + mock_response.aiter_bytes = AsyncMock() + mock_response.aclose = AsyncMock() + + mock_client.send = AsyncMock(return_value=mock_response) + mock_client.build_request = AsyncMock() + mock_client.aclose = AsyncMock() + + try: + response = await async_client.post( + "/v1/chat/completions", + headers={"Authorization": f"Bearer sk-{key.hashed_key}"}, + json={ + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}] + } + ) + + # Should succeed because balance (10,000 sats) > max_cost (5000 sats) + assert response.status_code == 200 + + finally: + MODELS.clear() + MODELS.extend(original_models) \ No newline at end of file