"""混合检索器测试.""" 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"] == "第一章"