Files
RAG-CUT/backend/tests/test_parent_child.py
2026-07-16 11:12:17 +08:00

81 lines
2.9 KiB
Python

"""Tests for parent/child delimiter splitting."""
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from rag_cut.models import Block, BlockType, SplitConfig, SplitMode
from rag_cut.pipeline import chunk_document
from rag_cut.splitters.parent_child import split_by_parent_child
def p(text: str) -> Block:
return Block(type=BlockType.PARAGRAPH, text=text)
class ParentChildSplitTest(unittest.TestCase):
def test_parent_and_child_chunks_are_linked(self) -> None:
blocks = [p("Intro A##Detail A1###Detail A2##Intro B###Detail B1")]
groups = split_by_parent_child(
blocks,
SplitConfig(
mode=SplitMode.PARENT_CHILD,
parent_delimiter="##",
child_delimiter="###",
max_chunk_size=2000,
child_max_size=500,
overlap=0,
),
)
parents = [g for g in groups if g.meta.get("is_section_parent")]
children = [g for g in groups if g.meta.get("is_sub_chunk")]
self.assertGreaterEqual(len(parents), 2)
self.assertGreaterEqual(len(children), 2)
self.assertTrue(all(g.meta.get("retrieval") is False for g in parents))
self.assertTrue(all(g.meta.get("retrieval") is True for g in children))
self.assertTrue(all(g.meta.get("parent_section_id") for g in children))
def test_requires_parent_delimiter(self) -> None:
with self.assertRaises(ValueError):
split_by_parent_child(
[p("hello")],
SplitConfig(mode=SplitMode.PARENT_CHILD, parent_delimiter=None),
)
def test_pipeline_persists_parent_chunk_ids(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
source = Path(tmp) / "pc.md"
source.write_text(
"Alpha parent##A1 child###A2 child##Beta parent###B1 child",
encoding="utf-8",
)
result = chunk_document(
source,
config=SplitConfig(
mode=SplitMode.PARENT_CHILD,
parent_delimiter="##",
child_delimiter="###",
max_chunk_size=2000,
child_max_size=400,
overlap=0,
),
storage_root=Path(tmp) / "storage",
)
self.assertEqual(result.split_mode, SplitMode.PARENT_CHILD)
self.assertEqual(result.split_config.get("chunk_strategy"), "parent_child_delimiter")
parents = [c for c in result.chunks if c.meta.get("is_section_parent")]
children = [c for c in result.chunks if c.meta.get("is_sub_chunk")]
self.assertTrue(parents)
self.assertTrue(children)
for child in children:
self.assertIn("parent_chunk_id", child.meta)
self.assertIsInstance(child.meta["parent_chunk_id"], int)
if __name__ == "__main__":
unittest.main()