"""RRF 融合与截断的单元测试""" import pytest from qdrant_client import models from app.core.ranker import finalize, rrf_fuse def _point(pid: str, score: float = 1.0) -> models.ScoredPoint: return models.ScoredPoint(id=pid, version=0, score=score, payload={}, vector=None) class TestRrfFuse: """rrf_fuse 的排序 / 去重 / 空输入 / k 参数行为""" def test_empty_input(self): """空列表输入返回 []""" assert rrf_fuse([]) == [] assert rrf_fuse([[], []]) == [] def test_single_list_fused_scores(self): """单路结果按 rank 计算 1/(k+rank) 融合分并降序返回""" fused = rrf_fuse([[_point("a"), _point("b"), _point("c")]], k=60) assert [p.id for p in fused] == ["a", "b", "c"] assert fused[0].score == pytest.approx(1 / 61) assert fused[1].score == pytest.approx(1 / 62) assert fused[2].score == pytest.approx(1 / 63) def test_dedup_merges_scores(self): """同一 point id 出现在多路结果中,只保留一条且分数累加""" list1 = [_point("a"), _point("b")] list2 = [_point("b"), _point("c")] fused = rrf_fuse([list1, list2], k=60) # b 在两路中分别 rank2/rank1,融合分最高;a 与 c 各 1/61 assert [p.id for p in fused] == ["b", "a", "c"] assert fused[0].score == pytest.approx(1 / 62 + 1 / 61) assert fused[1].score == pytest.approx(1 / 61) assert fused[2].score == pytest.approx(1 / 62) def test_k_affects_ranking(self): """k 参数改变融合权重,可改变最终排序 A 排名 [1, 10],B 排名 [5, 5]: - k=60 时 B 总分更高(两路均衡占优) - k=1 时 A 总分更高(头部 rank 权重被放大) """ l1 = [_point("a"), *[_point(f"n{i}") for i in range(3)], _point("b"), *[_point(f"m{i}") for i in range(5)]] l2 = [*[_point(f"x{i}") for i in range(4)], _point("b"), *[_point(f"y{i}") for i in range(4)], _point("a")] # l1:a rank1,b rank5;l2:b rank5,a rank10 fused_k60 = rrf_fuse([l1, l2], k=60) assert fused_k60[0].id == "b" fused_k1 = rrf_fuse([l1, l2], k=1) assert fused_k1[0].id == "a" def test_result_is_new_objects(self): """返回的是写回融合分的新对象,不修改原 point""" original = _point("a", score=0.99) fused = rrf_fuse([[original]], k=60) assert fused[0].score != 0.99 assert original.score == 0.99 class TestFinalize: """finalize 截断行为""" def test_truncate(self): points = [_point(str(i)) for i in range(10)] result = finalize(points, 5) assert [p.id for p in result] == ["0", "1", "2", "3", "4"] def test_truncate_larger_than_input(self): """final_k 大于列表长度时返回全部""" points = [_point("a"), _point("b")] assert finalize(points, 5) == points def test_empty(self): assert finalize([], 5) == []