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