"""嵌入服务单元测试(mock httpx/openai,不发起真实网络请求)""" from unittest.mock import AsyncMock, MagicMock, patch import httpx from app.config import settings from app.core.embeddings import ( EmbeddingService, LocalEmbeddingService, OpenAIEmbeddingService, create_embedding_service, ) def _fake_openai_response(dim: int, n: int) -> MagicMock: """构造 OpenAI embeddings.create 的假响应""" resp = MagicMock() resp.data = [MagicMock(embedding=[0.1 * (i + 1)] * dim) for i in range(n)] return resp class _FakeResponse: """httpx.Response 替代品""" def __init__(self, data: dict) -> None: self._data = data def raise_for_status(self) -> None: pass def json(self) -> dict: return self._data class _FakeAsyncClient: """httpx.AsyncClient 替代品,记录请求参数""" def __init__(self, data: dict, captured: dict, **kwargs) -> None: self._data = data self._captured = captured captured["timeout"] = kwargs.get("timeout") async def __aenter__(self) -> "_FakeAsyncClient": return self async def __aexit__(self, *args) -> bool: return False async def post(self, url: str, json: dict | None = None) -> _FakeResponse: self._captured["url"] = url self._captured["json"] = json return _FakeResponse(self._data) class TestOpenAIEmbeddingService: """OpenAI provider 行为""" async def test_embed_batch(self): """批量嵌入:模型与输入透传,返回与输入等长的向量""" with patch("app.core.embeddings.AsyncOpenAI") as mock_cls: client = mock_cls.return_value client.embeddings.create = AsyncMock(return_value=_fake_openai_response(dim=1536, n=2)) service = OpenAIEmbeddingService(api_key="k", base_url="http://x/v1", model="m") vectors = await service.embed(["你好", "world"]) assert len(vectors) == 2 assert all(len(v) == 1536 for v in vectors) client.embeddings.create.assert_awaited_once_with(model="m", input=["你好", "world"]) async def test_embed_empty(self): """空列表输入直接返回空列表,不调用 API""" with patch("app.core.embeddings.AsyncOpenAI") as mock_cls: client = mock_cls.return_value client.embeddings.create = AsyncMock() service = OpenAIEmbeddingService(api_key="k", base_url="http://x/v1", model="m") assert await service.embed([]) == [] client.embeddings.create.assert_not_called() async def test_dimension_mismatch_warns(self): """返回维度与配置不一致时记录 warning,不抛错""" with ( patch("app.core.embeddings.AsyncOpenAI") as mock_cls, patch("app.core.embeddings.logger") as mock_logger, ): client = mock_cls.return_value client.embeddings.create = AsyncMock(return_value=_fake_openai_response(dim=8, n=1)) service = OpenAIEmbeddingService(api_key="k", base_url="http://x/v1", model="m") vectors = await service.embed(["a"]) assert len(vectors[0]) == 8 mock_logger.warning.assert_called_once() class TestLocalEmbeddingService: """Ollama local provider 行为""" async def test_embed_batch(self, monkeypatch): """批量嵌入:POST /api/embed,payload 与 URL 正确,timeout 60s""" captured: dict = {} data = {"embeddings": [[0.1, 0.2], [0.3, 0.4]]} monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: _FakeAsyncClient(data, captured, **kw)) service = LocalEmbeddingService(base_url="http://localhost:11434/", model="bge-m3") vectors = await service.embed(["你好", "world"]) assert vectors == [[0.1, 0.2], [0.3, 0.4]] assert captured["url"] == "http://localhost:11434/api/embed" assert captured["json"] == {"model": "bge-m3", "input": ["你好", "world"]} assert captured["timeout"] == 60.0 async def test_embed_empty(self, monkeypatch): """空列表输入直接返回空列表,不发起请求""" called = [] monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: called.append(kw)) service = LocalEmbeddingService(base_url="http://localhost:11434", model="bge-m3") assert await service.embed([]) == [] assert called == [] async def test_dimension_mismatch_warns(self, monkeypatch): """返回维度与配置不一致时记录 warning,不抛错""" data = {"embeddings": [[0.1] * 8]} monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: _FakeAsyncClient(data, {}, **kw)) with patch("app.core.embeddings.logger") as mock_logger: service = LocalEmbeddingService(base_url="http://localhost:11434", model="bge-m3") vectors = await service.embed(["a"]) assert len(vectors[0]) == 8 mock_logger.warning.assert_called_once() class TestCreateEmbeddingService: """工厂函数选择逻辑""" def test_provider_openai(self, monkeypatch): monkeypatch.setattr(settings, "embedding_provider", "openai") monkeypatch.setattr(settings, "openai_api_key", "test-key") service = create_embedding_service() assert isinstance(service, OpenAIEmbeddingService) assert isinstance(service, EmbeddingService) def test_provider_local(self, monkeypatch): monkeypatch.setattr(settings, "embedding_provider", "local") monkeypatch.setattr(settings, "ollama_embedding_model", "bge-m3") service = create_embedding_service() assert isinstance(service, LocalEmbeddingService) assert service.model == "bge-m3" assert isinstance(service, EmbeddingService)