mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
ai generated tests (claude4opus)
This commit is contained in:
+1
-1
@@ -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"]
|
||||
|
||||
+17
@@ -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
|
||||
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user