"""重排序器测试.""" import pytest from src.core.reranker import Reranker class FakeCrossEncoder: """模拟 Cross-Encoder 模型.""" def predict(self, pairs, **kwargs): # 包含"重要"的 pair 分数高 scores = [] for pair in pairs: score = 5.0 if "重要" in pair[1] else 1.0 scores.append(score) return scores def test_reranker_returns_same_count(): """重排序不改变结果数量.""" reranker = Reranker(model_name="test-model") reranker._model = FakeCrossEncoder() candidates = [ {"content": "普通文档", "score": 0.8}, {"content": "重要文档", "score": 0.6}, {"content": "另一个普通", "score": 0.7}, ] result = reranker.rerank("查询", candidates, top_k=3) assert len(result) == 3 def test_reranker_promotes_relevant(): """重排序将更相关的内容提前.""" reranker = Reranker(model_name="test-model") reranker._model = FakeCrossEncoder() candidates = [ {"content": "普通 A", "score": 0.9}, {"content": "重要内容在这里", "score": 0.5}, {"content": "普通 B", "score": 0.7}, ] result = reranker.rerank("查询", candidates, top_k=3) assert "重要" in result[0]["content"] def test_reranker_truncates_to_top_k(): """rerank 截断到指定的 top_k.""" reranker = Reranker(model_name="test-model") reranker._model = FakeCrossEncoder() candidates = [ {"content": f"文档{i}", "score": 0.9 - i * 0.1} for i in range(20) ] result = reranker.rerank("查询", candidates, top_k=5) assert len(result) == 5 def test_reranker_empty_input(): """空输入返回空列表.""" reranker = Reranker(model_name="test-model") result = reranker.rerank("查询", [], top_k=5) assert result == [] def test_reranker_preserves_metadata(): """重排序保留文档元数据.""" reranker = Reranker(model_name="test-model") reranker._model = FakeCrossEncoder() candidates = [ { "content": "带元数据的文档", "score": 0.5, "source_file": "meta.md", "section_title": "第一章", } ] result = reranker.rerank("查询", candidates, top_k=1) assert result[0]["source_file"] == "meta.md" assert result[0]["section_title"] == "第一章" def test_reranker_score_replaced_with_rerank(): """重排序后 score 更新为 rerank_score.""" reranker = Reranker(model_name="test-model") reranker._model = FakeCrossEncoder() candidates = [{"content": "测试", "score": 0.5}] result = reranker.rerank("查询", candidates, top_k=1) assert "rerank_score" in result[0] assert result[0]["score"] == result[0]["rerank_score"]