"""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]