@@ -0,0 +1,51 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user