didi-lot1-ai/ai_platform/modules/embeddings/tests/conftest.py

82 lines
2.3 KiB
Python

"""Pytest fixtures for embeddings tests."""
from unittest.mock import AsyncMock
import pytest
from fastapi.testclient import TestClient
from embeddings.api.app import create_app
from embeddings.api.dependencies import init_concurrency_limiter
from embeddings.client import EmbeddingClient
from embeddings.config import EmbeddingSettings, SettingsCache
from embeddings.schemas import EmbeddingResponse
from embeddings.types import BackendType, EmbeddingData, EmbeddingUsage
@pytest.fixture(autouse=True)
def clear_settings_cache() -> None:
"""Clear settings cache before each test."""
SettingsCache.clear()
@pytest.fixture
def test_settings() -> EmbeddingSettings:
"""Create test settings with mocked values."""
return EmbeddingSettings(
default_backend="vllm",
enable_vllm=True,
enable_llamacpp=False,
external_url="http://localhost:54100",
host="127.0.0.1",
port=54100,
api_tokens=None,
)
@pytest.fixture
def mock_embedding_response() -> EmbeddingResponse:
"""Create a mock embedding response."""
return EmbeddingResponse(
data=[
EmbeddingData(
index=0,
embedding=[0.1, 0.2, 0.3, 0.4, 0.5],
)
],
model="test-model",
usage=EmbeddingUsage(prompt_tokens=5, total_tokens=5),
backend="vllm",
)
@pytest.fixture
def client_with_mock_backend(
test_settings: EmbeddingSettings, mock_embedding_response: EmbeddingResponse
) -> EmbeddingClient:
"""Create an EmbeddingClient with mocked backend."""
SettingsCache.set(test_settings)
client = EmbeddingClient(settings=test_settings)
# Mock the backend's embed method
backend = client.registry.get(BackendType.VLLM)
embeddings = [d.embedding for d in mock_embedding_response.data]
usage = mock_embedding_response.usage
backend.embed = AsyncMock(return_value=(embeddings, usage))
return client
@pytest.fixture
def app_client(test_settings: EmbeddingSettings) -> TestClient:
"""Create a FastAPI TestClient with test settings."""
SettingsCache.set(test_settings)
init_concurrency_limiter(test_settings.max_concurrent_requests)
app = create_app()
app.state.settings = test_settings
app.state.client = EmbeddingClient(settings=test_settings)
return TestClient(app)