Files
QMDSearch/tests/test_embeddings.py
kplam 51dc8dc4f6 Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈
- L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合)
- 文档三级总结与 2.5 级回退
- query 解析路由与分类
- /admin 管理页面
2026-07-29 21:24:40 +08:00

151 lines
5.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""嵌入服务单元测试(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/embedpayload 与 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)