Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈 - L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合) - 文档三级总结与 2.5 级回退 - query 解析路由与分类 - /admin 管理页面
This commit is contained in:
@@ -0,0 +1,150 @@
|
||||
"""嵌入服务单元测试(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)
|
||||
Reference in New Issue
Block a user