52 lines
1.6 KiB
Python
52 lines
1.6 KiB
Python
"""Tests for recall ranking and retrieval eligibility."""
|
|||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import unittest
|
||
|
|
|
||
|
|
from rag_cut.models import Chunk
|
||
|
|
from rag_cut.retrieval import recall_chunks
|
||
|
|
|
||
|
|
|
||
|
|
def chunk(index: int, content: str, *, retrieval: bool = True, heading: str = "") -> Chunk:
|
||
|
|
return Chunk(
|
||
|
|
index=index,
|
||
|
|
content=content,
|
||
|
|
char_count=len(content),
|
||
|
|
block_types=["paragraph"],
|
||
|
|
meta={"retrieval": retrieval, "heading": heading},
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
class RecallChunksTest(unittest.TestCase):
|
||
|
|
def test_relevant_chunk_ranks_first(self) -> None:
|
||
|
|
results, candidate_count = recall_chunks(
|
||
|
|
"股票交易密码",
|
||
|
|
[
|
||
|
|
chunk(0, "登录后可以修改股票交易密码", heading="账户安全"),
|
||
|
|
chunk(1, "年度报告及公司治理", heading="公司资料"),
|
||
|
|
],
|
||
|
|
top_k=2,
|
||
|
|
)
|
||
|
|
|
||
|
|
self.assertEqual(candidate_count, 2)
|
||
|
|
self.assertEqual(results[0]["chunk_index"], 0)
|
||
|
|
self.assertGreater(results[0]["score"], 0)
|
||
|
|
|
||
|
|
def test_preview_only_chunks_are_not_candidates(self) -> None:
|
||
|
|
results, candidate_count = recall_chunks(
|
||
|
|
"logo",
|
||
|
|
[chunk(0, "", retrieval=False), chunk(1, "Useful body")],
|
||
|
|
)
|
||
|
|
|
||
|
|
self.assertEqual(candidate_count, 1)
|
||
|
|
self.assertEqual(results, [])
|
||
|
|
|
||
|
|
def test_top_k_is_respected(self) -> None:
|
||
|
|
results, _ = recall_chunks("account", [chunk(i, f"account details {i}") for i in range(5)], top_k=2)
|
||
|
|
self.assertEqual(len(results), 2)
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
unittest.main()
|