ai generated tests (claude4opus)

This commit is contained in:
shroominic
2025-05-28 11:27:48 +00:00
parent 796ea1f7ae
commit 5be70738db
9 changed files with 1077 additions and 1 deletions
+1 -1
View File
@@ -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
View File
@@ -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
+63
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
+181
View File
@@ -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
+217
View File
@@ -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
+63
View File
@@ -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
+200
View File
@@ -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
+334
View File
@@ -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)