Files

162 lines
5.0 KiB
Python

"""混合检索器测试."""
import hashlib
import pytest
from src.core.retriever import HybridRetriever
class FakeEmbedder:
"""模拟嵌入器 — 返回伪向量."""
@property
def dimension(self) -> int:
return 4
def embed(self, texts: list[str]) -> list[list[float]]:
result = []
for t in texts:
h = hashlib.md5(t.encode()).digest()
vec = [float(b) / 255.0 for b in h[:4]]
result.append(vec)
return result
class FakeCollection:
"""模拟 ChromaDB collection."""
def __init__(self):
self._docs: list[dict] = []
def add(self, ids, embeddings, documents, metadatas):
for i, doc_id in enumerate(ids):
self._docs.append({
"id": doc_id,
"embedding": embeddings[i] if embeddings else [],
"document": documents[i],
"metadata": metadatas[i] if metadatas else {},
})
def query(self, query_embeddings, n_results, where=None, include=None):
# 返回所有已入库文档(模拟向量检索)
n = min(n_results, len(self._docs))
if n == 0:
return {"ids": [[]], "documents": [[]], "metadatas": [[]], "distances": [[]]}
ids_list = [d["id"] for d in self._docs[:n]]
docs_list = [d["document"] for d in self._docs[:n]]
metas_list = [d["metadata"] for d in self._docs[:n]]
dists_list = [0.2 + i * 0.05 for i in range(n)] # 伪距离
return {
"ids": [ids_list],
"documents": [docs_list],
"metadatas": [metas_list],
"distances": [dists_list],
}
def get(self, include=None):
return {
"ids": [d["id"] for d in self._docs],
"documents": [d["document"] for d in self._docs],
"metadatas": [d["metadata"] for d in self._docs],
}
def count(self) -> int:
return len(self._docs)
def delete(self, ids):
self._docs = [d for d in self._docs if d["id"] not in ids]
class FakeDB:
"""模拟 VectorDB."""
def __init__(self):
self._collections: dict[str, FakeCollection] = {}
def get_or_create_collection(self, name: str):
if name not in self._collections:
self._collections[name] = FakeCollection()
return self._collections[name]
@pytest.fixture
def retriever():
db = FakeDB()
embedder = FakeEmbedder()
return HybridRetriever(db, embedder, "test", bm25_weight=0.3)
def test_search_returns_list(retriever):
"""基本语义检索返回列表."""
results = retriever.search("测试查询", top_k=5)
assert isinstance(results, list)
def test_search_with_source_filter(retriever):
"""按 source_file 过滤."""
results = retriever.search("查询", top_k=5, source_file="doc.md")
assert isinstance(results, list)
def test_search_top_k_bounds(retriever):
"""top_k 在合理范围内."""
for k in [1, 10, 50]:
results = retriever.search("test", top_k=k)
assert len(results) <= k
def test_bm25_index_built_from_collection(retriever):
"""BM25 索引从 collection 文档构建."""
coll = retriever._db.get_or_create_collection("test")
coll.add(
ids=["doc_0", "doc_1", "doc_2"],
embeddings=[[0.1] * 4, [0.2] * 4, [0.3] * 4],
documents=["Python 是一门编程语言", "Java 也是编程语言", "今天天气很好"],
metadatas=[
{"source_file": "a.md"},
{"source_file": "b.md"},
{"source_file": "c.md"},
],
)
retriever._bm25_index = None # 强制重建
results = retriever.search("编程语言", top_k=2)
assert len(results) == 2
assert any("编程语言" in r["content"] for r in results)
def test_hybrid_score_fusion(retriever):
"""混合分数融合:向量分 + BM25 分加权."""
coll = retriever._db.get_or_create_collection("test")
coll.add(
ids=["d0", "d1"],
embeddings=[[1.0] * 4, [0.5] * 4],
documents=["Docker 容器化部署指南", "Python 数据分析入门"],
metadatas=[{"source_file": "x.md"}, {"source_file": "y.md"}],
)
retriever._bm25_index = None
results = retriever.search("Docker 部署", top_k=2)
assert len(results) >= 1
assert "Docker" in results[0]["content"]
def test_empty_collection_returns_empty(retriever):
"""空 collection 返回空列表."""
results = retriever.search("查询", top_k=5)
assert results == []
def test_metadata_in_results(retriever):
"""结果中包含完整元数据."""
coll = retriever._db.get_or_create_collection("test")
coll.add(
ids=["meta_test"],
embeddings=[[0.5] * 4],
documents=["带元数据的文档"],
metadatas=[{"source_file": "meta.md", "section_title": "第一章", "heading_level": 1}],
)
retriever._bm25_index = None
results = retriever.search("元数据", top_k=1)
assert len(results) == 1
assert results[0]["source_file"] == "meta.md"
assert results[0]["section_title"] == "第一章"