81 lines
2.9 KiB
Python
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()
|