51dc8dc4f6
- FastAPI + Qdrant + Redis + Ollama 技术栈 - L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合) - 文档三级总结与 2.5 级回退 - query 解析路由与分类 - /admin 管理页面
151 lines
5.7 KiB
Python
151 lines
5.7 KiB
Python
"""嵌入服务单元测试(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)
|