"""Tests for chat completion routes.""" import pytest from unittest.mock import MagicMock from fastapi.testclient import TestClient from app.main import app from app.models import ClientKey, Provider @pytest.fixture def mock_client_key(): """Return a valid client key for testing.""" return ClientKey( id=1, key="test-api-key", name="test-client", is_active=True, created_at="2024-01-01" ) @pytest.fixture def mock_provider(): """Return a mock provider.""" return Provider( id=1, name="openai", base_url="https://api.openai.com/v1", api_type="openai", is_active=True ) class TestChatCompletions: """Test /v1/chat/completions endpoint.""" def test_chat_unauthorized(self): """Test chat without auth returns 401.""" client = TestClient(app) response = client.post( "/v1/chat/completions", json={"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]} ) assert response.status_code == 401 def test_chat_invalid_key(self): """Test chat with invalid key returns 401.""" client = TestClient(app) response = client.post( "/v1/chat/completions", json={"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}, headers={"Authorization": "Bearer invalid-key"} ) assert response.status_code == 401 def test_chat_unsupported_model(self, mock_client_key, mock_provider): """Test chat with unsupported model returns 400.""" from app.middleware.auth import get_current_client from app.services.router import model_router app.dependency_overrides[get_current_client] = lambda: mock_client_key async def mock_route_model(model): return (None, f"Unsupported model: {model}") original_route = model_router.route_model model_router.route_model = mock_route_model try: client = TestClient(app) response = client.post( "/v1/chat/completions", json={"model": "unknown-model", "messages": [{"role": "user", "content": "hi"}]}, headers={"Authorization": "Bearer test-api-key"} ) assert response.status_code == 400 finally: model_router.route_model = original_route app.dependency_overrides.clear() def test_chat_success(self, mock_client_key, mock_provider): """Test successful chat completion.""" from app.middleware.auth import get_current_client from app.services.router import model_router from app.services.key_selector import key_selector from app.services.logger import async_logger from app.providers.litellm_wrapper import LiteLLMWrapper app.dependency_overrides[get_current_client] = lambda: mock_client_key mock_response = MagicMock() mock_response.model_dump.return_value = { "id": "chatcmpl-test", "choices": [{"message": {"content": "Hello!"}}] } original_route = model_router.route_model original_get_key = key_selector.get_next_key original_chat = LiteLLMWrapper.chat_completion original_log = async_logger.log async def mock_route_model(model): return (mock_provider, None) async def mock_get_key(provider_id): return MagicMock(key="sk-test") async def mock_chat(*args, **kwargs): return (mock_response, 10, 20) async def mock_log(entry): pass model_router.route_model = mock_route_model key_selector.get_next_key = mock_get_key LiteLLMWrapper.chat_completion = staticmethod(lambda *a, **k: mock_chat(*a, **k)) async_logger.log = mock_log try: client = TestClient(app) response = client.post( "/v1/chat/completions", json={ "model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}], "temperature": 0.7 }, headers={"Authorization": "Bearer test-api-key"} ) assert response.status_code == 200 finally: model_router.route_model = original_route key_selector.get_next_key = original_get_key LiteLLMWrapper.chat_completion = original_chat async_logger.log = original_log app.dependency_overrides.clear() class TestListModels: """Test /v1/models endpoint.""" def test_list_models_unauthorized(self): """Test list models without auth.""" client = TestClient(app) response = client.get("/v1/models") assert response.status_code == 401 def test_list_models_success(self, mock_client_key): """Test list models with valid auth.""" from app.middleware.auth import get_current_client app.dependency_overrides[get_current_client] = lambda: mock_client_key try: client = TestClient(app) response = client.get( "/v1/models", headers={"Authorization": "Bearer test-api-key"} ) assert response.status_code == 200 data = response.json() assert data["object"] == "list" assert len(data["data"]) > 0 finally: app.dependency_overrides.clear()