82 lines
2.3 KiB
Python
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)
|