211 lines
7.1 KiB
Python
211 lines
7.1 KiB
Python
"""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]
|