feat: 接入 BGE-M3 本地嵌入默认配置 + 新增 Qwen3-Reranker 重排

- .env.example 嵌入默认改为本地 BGE-M3(EMBEDDING_PROVIDER=local, dim=1024),NAS 全新部署即用
- 新增 app/services/reranker.py:基于 Ollama /api/rerank 的 cross-encoder 重排服务
- retriever 在 chunk 候选阶段接入语义精排,调用失败优雅降级为 RRF 顺序
- docker-compose 拉取 qwen3-reranker:0.6b,OLLAMA_MAX_LOADED_MODELS 提至 3
- .gitignore 排除 .workbuddy/ 与 .trae/(防止误提交项目记忆与 IDE spec)
- 新增 reranker 单元与集成测试,全量测试 529 passed
This commit is contained in:
2026-08-01 00:08:38 +08:00
parent 6c6f690788
commit 6ae91679e2
8 changed files with 320 additions and 9 deletions
+128
View File
@@ -0,0 +1,128 @@
"""RerankerService 单元测试(mock httpx,不发起真实网络请求)"""
from typing import Any
import httpx
import pytest
from app.config import settings
from app.services.reranker import OllamaRerankerService, create_reranker_service
def _make_response(status_code: int, data: dict | None = None) -> httpx.Response:
"""构造带 request 上下文的 httpx.Responseraise_for_status 依赖 request"""
request = httpx.Request("POST", "http://localhost:11434/api/rerank")
if data is None:
return httpx.Response(status_code, request=request)
return httpx.Response(status_code, json=data, request=request)
class _FakeAsyncClient:
"""httpx.AsyncClient 替代品:记录请求参数,返回预设响应或抛出预设异常"""
def __init__(
self,
captured: dict,
response: httpx.Response | None = None,
error: Exception | None = None,
**kwargs: Any,
) -> None:
self._captured = captured
self._response = response
self._error = error
captured["timeout"] = kwargs.get("timeout")
async def __aenter__(self) -> "_FakeAsyncClient":
return self
async def __aexit__(self, *args: object) -> bool:
return False
async def post(self, url: str, json: dict | None = None) -> httpx.Response:
self._captured["url"] = url
self._captured["json"] = json
if self._error is not None:
raise self._error
assert self._response is not None
return self._response
def _client(monkeypatch: pytest.MonkeyPatch, captured: dict, **kw: Any) -> None:
monkeypatch.setattr(httpx, "AsyncClient", lambda **k: _FakeAsyncClient(captured, **k, **kw))
class TestRerank:
async def test_returns_scores_aligned_to_documents(self, monkeypatch: pytest.MonkeyPatch) -> None:
captured: dict = {}
_client(
monkeypatch,
captured,
response=_make_response(
200,
{
"results": [
{"index": 0, "relevance_score": 0.2},
{"index": 1, "relevance_score": 0.9},
{"index": 2, "relevance_score": 0.5},
]
},
),
)
service = OllamaRerankerService(base_url="http://localhost:11434/", model="qwen3-reranker:0.6b")
scores = await service.rerank("查询", ["doc-a", "doc-b", "doc-c"])
assert scores == [0.2, 0.9, 0.5]
assert captured["url"] == "http://localhost:11434/api/rerank"
assert captured["json"] == {
"model": "qwen3-reranker:0.6b",
"query": "查询",
"documents": ["doc-a", "doc-b", "doc-c"],
}
async def test_empty_documents_returns_empty(self, monkeypatch: pytest.MonkeyPatch) -> None:
captured: dict = {}
_client(monkeypatch, captured, response=_make_response(200, {"results": []}))
service = OllamaRerankerService(base_url="http://localhost:11434", model="qwen3-reranker:0.6b")
scores = await service.rerank("查询", [])
assert scores == []
# 空文档不发起请求
assert "url" not in captured
async def test_partial_results_zero_fill_missing(self, monkeypatch: pytest.MonkeyPatch) -> None:
captured: dict = {}
# 仅返回 index 1index 0/2 缺失 → 补 0
_client(
monkeypatch,
captured,
response=_make_response(200, {"results": [{"index": 1, "relevance_score": 0.7}]}),
)
service = OllamaRerankerService(base_url="http://localhost:11434", model="qwen3-reranker:0.6b")
scores = await service.rerank("查询", ["a", "b", "c"])
assert scores == [0.0, 0.7, 0.0]
async def test_http_error_raises(self, monkeypatch: pytest.MonkeyPatch) -> None:
_client(monkeypatch, {}, response=_make_response(500))
service = OllamaRerankerService(base_url="http://localhost:11434", model="qwen3-reranker:0.6b")
with pytest.raises(httpx.HTTPStatusError):
await service.rerank("查询", ["a", "b"])
class TestCreateRerankerService:
def test_disabled_returns_none(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(settings, "reranker_enabled", False)
assert create_reranker_service() is None
def test_enabled_returns_service(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(settings, "reranker_enabled", True)
monkeypatch.setattr(settings, "reranker_model", "qwen3-reranker:0.6b")
monkeypatch.setattr(settings, "ollama_base_url", "http://ollama:11434")
monkeypatch.setattr(settings, "reranker_timeout", 30.0)
service = create_reranker_service()
assert isinstance(service, OllamaRerankerService)
assert service.model == "qwen3-reranker:0.6b"
+64 -1
View File
@@ -82,15 +82,29 @@ def _route(
)
def _make_retriever(qdrant: FakeQdrant, route: RouteDecision) -> Retriever:
def _make_retriever(qdrant: FakeQdrant, route: RouteDecision, reranker: Any = None) -> Retriever:
return Retriever(
qdrant=qdrant, # type: ignore[arg-type]
query_parser=FakeParser(route), # type: ignore[arg-type]
embedding=FakeEmbedding(),
sparse_encoder=SparseEncoder(),
reranker=reranker,
)
class FakeReranker:
"""假 RerankerService:返回预设分数或抛出预设异常"""
def __init__(self, scores: list[float] | None = None, exc: Exception | None = None) -> None:
self.scores = scores if scores is not None else []
self.exc = exc
async def rerank(self, query: str, documents: list[str]) -> list[float]:
if self.exc is not None:
raise self.exc
return self.scores
def _calls(qdrant: FakeQdrant, collection: str) -> list[dict[str, Any]]:
return [c for c in qdrant.calls if c["collection"] == collection]
@@ -348,3 +362,52 @@ class TestSearchApi:
resp = client.get("/api/v1/health")
assert resp.status_code == 200
assert resp.json() == {"status": "ok"}
class TestReranker:
"""重排启用时对 chunk 候选精排;关闭或失败则退化为 RRF 顺序"""
@staticmethod
def _full_qdrant() -> FakeQdrant:
return FakeQdrant(
{
COLLECTION_L1: [[_point("l1a", "d1")]],
COLLECTION_L2: [[_point("l2a", "d1", "章节A")]],
COLLECTION_L3: [[_point("l3a", "d1", "章节A")]],
COLLECTION_CHUNKS: [
[
_point("c1", "d1", "章节A"),
_point("c2", "d1", "章节A"),
_point("c3", "d1", "章节A"),
]
],
}
)
async def test_enabled_reorders_by_relevance(self):
"""重排按相关性分数降序重排,并将 score 写回相关性分数"""
reranker = FakeReranker(scores=[0.1, 0.9, 0.5]) # c1<c3<c2
retriever = _make_retriever(self._full_qdrant(), _route(), reranker=reranker)
resp = await retriever.search(SearchRequest(query="测试查询"))
# 重排后顺序为 c2, c3, c1
assert [h.text for h in resp.hits] == ["text-c2", "text-c3", "text-c1"]
assert resp.hits[0].score == 0.9
async def test_disabled_keeps_rrf_order(self):
"""重排未注入(关闭)时保持 RRF 融合原始顺序"""
retriever = _make_retriever(self._full_qdrant(), _route(), reranker=None)
resp = await retriever.search(SearchRequest(query="测试查询"))
assert [h.text for h in resp.hits] == ["text-c1", "text-c2", "text-c3"]
async def test_failure_falls_back_to_rrf(self):
"""重排调用异常时退化为 RRF 顺序,检索不中断"""
reranker = FakeReranker(exc=RuntimeError("ollama 不可用"))
retriever = _make_retriever(self._full_qdrant(), _route(), reranker=reranker)
resp = await retriever.search(SearchRequest(query="测试查询"))
assert [h.text for h in resp.hits] == ["text-c1", "text-c2", "text-c3"]