- Add tests for all backend modules: config, database, models, services, routes, middleware, and providers - Separate dev dependencies using dependency-groups - Update CLAUDE.md with test commands and project structure - Add .coverage and .pytest_cache to gitignore 88 tests covering: - Authentication and authorization - Model routing and key selection - Async logging with batch processing - All API endpoints (chat, openai, anthropic, health) - LiteLLM wrapper (mocked external calls) Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
98 lines
4.1 KiB
Python
98 lines
4.1 KiB
Python
"""Tests for model router service."""
|
|
import pytest
|
|
from unittest.mock import patch
|
|
from app.services.router import ModelRouter, model_router
|
|
from app.models import Provider
|
|
|
|
|
|
class TestModelRouter:
|
|
"""Test ModelRouter."""
|
|
|
|
def test_get_provider_name_gpt(self):
|
|
"""Test getting provider for GPT models."""
|
|
assert ModelRouter.get_provider_name("gpt-4o") == "openai"
|
|
assert ModelRouter.get_provider_name("gpt-4-turbo") == "openai"
|
|
assert ModelRouter.get_provider_name("gpt-3.5-turbo") == "openai"
|
|
|
|
def test_get_provider_name_o1(self):
|
|
"""Test getting provider for o1 models."""
|
|
assert ModelRouter.get_provider_name("o1") == "openai"
|
|
assert ModelRouter.get_provider_name("o1-mini") == "openai"
|
|
assert ModelRouter.get_provider_name("o1-preview") == "openai"
|
|
|
|
def test_get_provider_name_o3(self):
|
|
"""Test getting provider for o3 models."""
|
|
assert ModelRouter.get_provider_name("o3-mini") == "openai"
|
|
|
|
def test_get_provider_name_claude(self):
|
|
"""Test getting provider for Claude models."""
|
|
assert ModelRouter.get_provider_name("claude-3-opus") == "anthropic"
|
|
assert ModelRouter.get_provider_name("claude-3-sonnet") == "anthropic"
|
|
assert ModelRouter.get_provider_name("claude-3-haiku") == "anthropic"
|
|
|
|
def test_get_provider_name_unknown(self):
|
|
"""Test getting provider for unknown models."""
|
|
assert ModelRouter.get_provider_name("unknown-model") is None
|
|
assert ModelRouter.get_provider_name("llama-2") is None
|
|
assert ModelRouter.get_provider_name("mistral") is None
|
|
|
|
def test_get_provider_name_case_insensitive(self):
|
|
"""Test model name matching is case insensitive."""
|
|
assert ModelRouter.get_provider_name("GPT-4O") == "openai"
|
|
assert ModelRouter.get_provider_name("CLAUDE-3-OPUS") == "anthropic"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_provider_success(self, mock_db):
|
|
"""Test getting provider by name."""
|
|
with patch("app.services.router.db_connection", mock_db):
|
|
provider = await ModelRouter.get_provider("openai")
|
|
assert provider is not None
|
|
assert isinstance(provider, Provider)
|
|
assert provider.name == "openai"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_provider_anthropic(self, mock_db):
|
|
"""Test getting Anthropic provider."""
|
|
with patch("app.services.router.db_connection", mock_db):
|
|
provider = await ModelRouter.get_provider("anthropic")
|
|
assert provider is not None
|
|
assert provider.name == "anthropic"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_provider_not_found(self, mock_db):
|
|
"""Test getting non-existent provider."""
|
|
with patch("app.services.router.db_connection", mock_db):
|
|
provider = await ModelRouter.get_provider("unknown")
|
|
assert provider is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_route_model_success(self, mock_db):
|
|
"""Test routing model to provider."""
|
|
with patch("app.services.router.db_connection", mock_db):
|
|
provider, error = await ModelRouter.route_model("gpt-4o")
|
|
assert provider is not None
|
|
assert error is None
|
|
assert provider.name == "openai"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_route_model_anthropic(self, mock_db):
|
|
"""Test routing Claude model."""
|
|
with patch("app.services.router.db_connection", mock_db):
|
|
provider, error = await ModelRouter.route_model("claude-3-opus")
|
|
assert provider is not None
|
|
assert error is None
|
|
assert provider.name == "anthropic"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_route_model_unsupported(self, mock_db):
|
|
"""Test routing unsupported model."""
|
|
with patch("app.services.router.db_connection", mock_db):
|
|
provider, error = await ModelRouter.route_model("llama-2")
|
|
assert provider is None
|
|
assert "Unsupported model" in error
|
|
|
|
def test_global_instance_exists(self):
|
|
"""Test global model_router instance exists."""
|
|
assert model_router is not None
|
|
assert isinstance(model_router, ModelRouter)
|