Files
2026-07-16 11:12:17 +08:00

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, "![logo](logo.png)", 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()