Files
2026-07-16 11:12:17 +08:00

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]