@@ -0,0 +1,210 @@
|
||||
"""Default structure-aware splitting."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from rag_cut.models import Block, BlockType, SplitConfig
|
||||
from rag_cut.splitters.heading_splitter import (
|
||||
chunk_groups_to_block_groups,
|
||||
has_meaningful_headings,
|
||||
split_by_heading_hierarchy,
|
||||
)
|
||||
|
||||
ATOMIC_TYPES = {BlockType.TABLE, BlockType.IMAGE}
|
||||
|
||||
|
||||
def _same_layout_context(prev: Block, curr: Block) -> bool:
|
||||
"""True when blocks should stay together to preserve image/text position."""
|
||||
if prev.meta.get("page") != curr.meta.get("page"):
|
||||
return False
|
||||
if prev.type == BlockType.PARAGRAPH and curr.type in {BlockType.IMAGE, BlockType.TABLE}:
|
||||
return True
|
||||
if prev.type == BlockType.HEADING and curr.type in {BlockType.PARAGRAPH, BlockType.IMAGE, BlockType.TABLE}:
|
||||
return prev.meta.get("parent_heading") == curr.meta.get("parent_heading") or not prev.meta.get("parent_heading")
|
||||
if prev.type == BlockType.IMAGE and curr.type == BlockType.PARAGRAPH:
|
||||
return True
|
||||
if prev.type == BlockType.TABLE and curr.type == BlockType.PARAGRAPH:
|
||||
return True
|
||||
if prev.type == BlockType.IMAGE and curr.type == BlockType.IMAGE:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _is_heading(block: Block) -> bool:
|
||||
return block.type == BlockType.HEADING
|
||||
|
||||
|
||||
def _meaningful_headings(blocks: list[Block]) -> bool:
|
||||
headings = [b for b in blocks if _is_heading(b)]
|
||||
if len(headings) < 2:
|
||||
return False
|
||||
substantial = [h for h in headings if len((h.text or "").strip()) >= 8]
|
||||
return len(substantial) >= 2
|
||||
|
||||
|
||||
def _split_by_headings(blocks: list[Block]) -> list[list[Block]]:
|
||||
"""Outline-aware: split when a heading of same-or-higher level appears."""
|
||||
if not any(_is_heading(b) for b in blocks):
|
||||
return []
|
||||
|
||||
sections: list[list[Block]] = []
|
||||
current: list[Block] = []
|
||||
stack: list[int] = []
|
||||
|
||||
for block in blocks:
|
||||
if _is_heading(block):
|
||||
level = block.level or 1
|
||||
while stack and stack[-1] >= level:
|
||||
stack.pop()
|
||||
if current:
|
||||
sections.append(current)
|
||||
current = []
|
||||
stack.append(level)
|
||||
current.append(block)
|
||||
else:
|
||||
current.append(block)
|
||||
|
||||
if current:
|
||||
sections.append(current)
|
||||
return sections
|
||||
|
||||
|
||||
def _split_by_page(blocks: list[Block], config: SplitConfig) -> list[list[Block]]:
|
||||
"""Prefer page boundaries for PDF/manual style documents."""
|
||||
if not any(b.meta.get("page") for b in blocks):
|
||||
return []
|
||||
|
||||
groups: list[list[Block]] = []
|
||||
current: list[Block] = []
|
||||
current_page: int | None = None
|
||||
|
||||
for block in blocks:
|
||||
page = block.meta.get("page")
|
||||
if current and page != current_page:
|
||||
groups.append(current)
|
||||
current = []
|
||||
current_page = page
|
||||
current.append(block)
|
||||
|
||||
if current:
|
||||
groups.append(current)
|
||||
|
||||
return _merge_oversized_sections(groups, config)
|
||||
|
||||
|
||||
def _split_by_size(blocks: list[Block], config: SplitConfig) -> list[list[Block]]:
|
||||
"""Fallback: pack blocks up to max_chunk_size without splitting atomic blocks."""
|
||||
groups: list[list[Block]] = []
|
||||
current: list[Block] = []
|
||||
current_len = 0
|
||||
|
||||
def flush() -> None:
|
||||
nonlocal current, current_len
|
||||
if current:
|
||||
groups.append(current)
|
||||
current = []
|
||||
current_len = 0
|
||||
|
||||
for block in blocks:
|
||||
rendered = block.render()
|
||||
block_len = len(rendered) + 2
|
||||
page = block.meta.get("page")
|
||||
|
||||
if block.type in ATOMIC_TYPES and current_len + block_len > config.max_chunk_size and current:
|
||||
if not _same_layout_context(current[-1], block):
|
||||
flush()
|
||||
|
||||
if block.type not in ATOMIC_TYPES and block_len > config.max_chunk_size:
|
||||
if current:
|
||||
flush()
|
||||
text = block.text or block.markdown
|
||||
start = 0
|
||||
while start < len(text):
|
||||
end = min(start + config.max_chunk_size, len(text))
|
||||
piece = Block(type=block.type, text=text[start:end], level=block.level, meta=block.meta)
|
||||
groups.append([piece])
|
||||
if end >= len(text):
|
||||
break
|
||||
start = max(end - config.overlap, start + 1)
|
||||
continue
|
||||
|
||||
if current_len + block_len > config.max_chunk_size and current:
|
||||
# Keep image with preceding heading/body on the same page
|
||||
if _same_layout_context(current[-1], block):
|
||||
pass
|
||||
else:
|
||||
flush()
|
||||
elif (
|
||||
current
|
||||
and page is not None
|
||||
and current[-1].meta.get("page") != page
|
||||
and current_len >= min(400, config.max_chunk_size // 3)
|
||||
):
|
||||
flush()
|
||||
|
||||
current.append(block)
|
||||
current_len += block_len
|
||||
|
||||
flush()
|
||||
return groups
|
||||
|
||||
|
||||
def _merge_oversized_sections(sections: list[list[Block]], config: SplitConfig) -> list[list[Block]]:
|
||||
result: list[list[Block]] = []
|
||||
for section in sections:
|
||||
rendered_len = sum(len(b.render()) + 2 for b in section)
|
||||
if rendered_len <= config.max_chunk_size:
|
||||
result.append(section)
|
||||
else:
|
||||
result.extend(_split_by_size(section, config))
|
||||
return result
|
||||
|
||||
|
||||
def split_default(blocks: list[Block], config: SplitConfig) -> list[list[Block]]:
|
||||
if not blocks:
|
||||
return []
|
||||
|
||||
# Spreadsheet: single table block — default = chunk by groups of rows with header
|
||||
if len(blocks) == 1 and blocks[0].type == BlockType.TABLE and blocks[0].meta.get("rows"):
|
||||
from rag_cut.splitters.by_row import split_by_row
|
||||
|
||||
row_config = SplitConfig(
|
||||
mode=config.mode,
|
||||
header_row_start=1,
|
||||
header_row_end=1,
|
||||
start_row=2,
|
||||
rows_per_chunk=max(1, min(10, len(blocks[0].meta["rows"]) // 5 or 1)),
|
||||
max_chunk_size=config.max_chunk_size,
|
||||
overlap=config.overlap,
|
||||
)
|
||||
return split_by_row(blocks, row_config)
|
||||
|
||||
# Heading hierarchy first: same section keeps body/images/tables/captions together
|
||||
if has_meaningful_headings(blocks):
|
||||
heading_groups = split_by_heading_hierarchy(blocks, config)
|
||||
if heading_groups:
|
||||
block_groups, _ = chunk_groups_to_block_groups(heading_groups)
|
||||
return block_groups
|
||||
|
||||
page_sections = _split_by_page(blocks, config)
|
||||
if len(page_sections) > 1:
|
||||
return page_sections
|
||||
|
||||
return _split_by_size(blocks, config)
|
||||
|
||||
|
||||
def split_default_with_meta(blocks: list[Block], config: SplitConfig) -> tuple[list[list[Block]], list[dict]]:
|
||||
"""Like split_default but also returns per-group metadata (heading sections)."""
|
||||
if not blocks:
|
||||
return [], []
|
||||
|
||||
if len(blocks) == 1 and blocks[0].type == BlockType.TABLE and blocks[0].meta.get("rows"):
|
||||
groups = split_default(blocks, config)
|
||||
return groups, [{} for _ in groups]
|
||||
|
||||
if has_meaningful_headings(blocks):
|
||||
heading_groups = split_by_heading_hierarchy(blocks, config)
|
||||
if heading_groups:
|
||||
return chunk_groups_to_block_groups(heading_groups)
|
||||
|
||||
groups = split_default(blocks, config)
|
||||
return groups, [{} for _ in groups]
|
||||
Reference in New Issue
Block a user