Files
RAG-CUT/backend/rag_cut/splitters/heading_splitter.py
T
2026-07-16 11:12:17 +08:00

345 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Heading-hierarchy-first document splitting."""
from __future__ import annotations
import re
from dataclasses import dataclass, field
from rag_cut.models import Block, BlockType, SplitConfig
# e.g. "1.2 ACCOUNT STATUS CODE MASTER", "10.上升三角形態" (space after '.' optional)
NUMBERED_HEADING_RE = re.compile(
r"^\s*(\d+(?:\.\d+)*)(?:\.|.)?\s*([A-Za-z0-9\u4e00-\u9fff][A-Za-z0-9\u4e00-\u9fff&/ \-_::]{2,})\s*$"
)
FIGURE_TABLE_RE = re.compile(
r"^\s*(?:Figure|Fig\.|图|表|Table)\s*[\d.]+",
re.I,
)
STEP_RE = re.compile(
r"^\s*(?:Step\s*\d+|步骤\s*\d+|\d{1,2}[..、]\s*(?:点击|點擊|选择|選擇|输入|輸入))",
re.I,
)
ATOMIC_TYPES = {BlockType.TABLE, BlockType.IMAGE}
MAX_HEADING_CHARS = 100
@dataclass
class ChunkGroup:
blocks: list[Block]
meta: dict = field(default_factory=dict)
def _text(block: Block) -> str:
return (block.text or block.markdown or "").strip()
def _rendered_len(blocks: list[Block]) -> int:
return sum(len(b.render()) + 2 for b in blocks)
def numbered_heading_level(text: str) -> int | None:
match = NUMBERED_HEADING_RE.match(text.strip())
if not match:
return None
return match.group(1).count(".") + 1
def infer_heading_level(block: Block) -> int:
if block.type == BlockType.HEADING and block.level:
numbered = numbered_heading_level(block.text)
if numbered:
return numbered
return block.level
numbered = numbered_heading_level(_text(block))
if numbered:
return numbered
return block.level or 1
def is_heading_block(block: Block) -> bool:
text = _text(block)
if not text or len(text) > MAX_HEADING_CHARS:
return False
if block.type == BlockType.HEADING:
return True
if NUMBERED_HEADING_RE.match(text):
return True
if block.meta.get("font_size") and block.meta.get("body_font_size"):
return block.meta["font_size"] >= block.meta["body_font_size"] + 1.5
return False
def normalize_heading_block(block: Block) -> Block:
text = _text(block)
if block.type == BlockType.HEADING and len(text) > MAX_HEADING_CHARS:
# Parser sometimes merges title+body then marks the blob as heading.
return Block(type=BlockType.PARAGRAPH, text=text, meta=dict(block.meta))
numbered = numbered_heading_level(text)
if block.type == BlockType.HEADING:
level = numbered or block.level or 1
return block.model_copy(update={"level": level})
if numbered and NUMBERED_HEADING_RE.match(text):
return Block(
type=BlockType.HEADING,
text=text,
level=numbered,
meta=dict(block.meta),
)
return block
def has_meaningful_headings(blocks: list[Block]) -> bool:
headings = [normalize_heading_block(b) for b in blocks]
count = sum(1 for b in headings if is_heading_block(b) or b.type == BlockType.HEADING)
return count >= 2
def _section_key(block: Block | None, fallback: int) -> str:
if block is None:
return f"section-{fallback}"
oi = block.meta.get("order_index", fallback)
title = re.sub(r"\W+", "-", (_text(block) or "heading"))[:48]
return f"h-{oi}-{title}"
@dataclass
class _OpenSection:
level: int
start_index: int
blocks: list[Block]
def _split_primary_sections(blocks: list[Block]) -> list[list[Block]]:
"""
Split so each heading owns its content until the next same-or-higher-level heading.
Example: 1.2 section runs until 1.3 (same level) or 2.0 (higher level).
Nested sub-headings (1.2.1) stay inside the 1.2 section.
"""
if not blocks:
return []
open_sections: list[_OpenSection] = []
finished: list[tuple[int, list[Block]]] = []
for raw in blocks:
block = normalize_heading_block(raw)
if is_heading_block(block) or block.type == BlockType.HEADING:
level = infer_heading_level(block)
while open_sections and open_sections[-1].level >= level:
sec = open_sections.pop()
finished.append((sec.start_index, sec.blocks))
start = int(block.meta.get("order_index", len(finished)))
open_sections.append(_OpenSection(level=level, start_index=start, blocks=[block]))
elif open_sections:
open_sections[-1].blocks.append(block)
while open_sections:
sec = open_sections.pop()
finished.append((sec.start_index, sec.blocks))
finished.sort(key=lambda item: item[0])
return [sec_blocks for _, sec_blocks in finished]
def _section_heading(section: list[Block]) -> Block | None:
for block in section:
nb = normalize_heading_block(block)
if nb.type == BlockType.HEADING or is_heading_block(nb):
return nb
return None
def _split_by_child_headings(section: list[Block], parent_level: int) -> list[list[Block]]:
"""Split an oversized section by deeper sub-headings."""
child_sections: list[list[Block]] = []
current: list[Block] = []
parent_heading = _section_heading(section)
for block in section:
nb = normalize_heading_block(block)
if (
block is not parent_heading
and (nb.type == BlockType.HEADING or is_heading_block(nb))
and infer_heading_level(nb) > parent_level
):
if current:
child_sections.append(current)
current = [block]
else:
current.append(block)
if current:
child_sections.append(current)
return child_sections if len(child_sections) > 1 else [section]
def _is_split_marker(block: Block) -> bool:
if block.type in ATOMIC_TYPES:
return False
text = _text(block)
if not text:
return False
return bool(FIGURE_TABLE_RE.match(text) or STEP_RE.match(text))
def _split_by_content_markers(section: list[Block], config: SplitConfig) -> list[list[Block]]:
"""Fallback: split at figure/table/step markers while keeping atomic blocks intact."""
if _rendered_len(section) <= config.max_chunk_size:
return [section]
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 section:
blen = len(block.render()) + 2
if (
current
and _is_split_marker(block)
and current_len >= min(500, config.max_chunk_size // 4)
and current_len + blen > config.max_chunk_size
):
flush()
elif current_len + blen > config.max_chunk_size and current:
if block.type in ATOMIC_TYPES:
flush()
elif not _is_split_marker(block):
flush()
current.append(block)
current_len += blen
flush()
return groups if groups else [section]
def _prepend_parent_heading(section: list[Block], parent: Block | None) -> list[Block]:
if not parent:
return section
parent_text = _text(parent)
if section and _text(normalize_heading_block(section[0])) == parent_text:
return section
return [parent] + section
def _split_oversized_section(
section: list[Block],
config: SplitConfig,
section_id: str,
) -> list[ChunkGroup]:
"""Split an oversized heading section by sub-headings/markers/size — no parent chunk."""
parent_heading = _section_heading(section)
parent_level = infer_heading_level(parent_heading) if parent_heading else 1
parent_title = _text(parent_heading) if parent_heading else ""
child_sections = _split_by_child_headings(section, parent_level)
if len(child_sections) == 1:
child_sections = _split_by_content_markers(section, config)
if len(child_sections) == 1:
from rag_cut.splitters.default_splitter import _split_by_size
child_sections = _split_by_size(section, config)
parts: list[ChunkGroup] = []
for idx, child in enumerate(child_sections):
child_heading = _section_heading(child)
blocks = _prepend_parent_heading(child, parent_heading)
parts.append(
ChunkGroup(
blocks=blocks,
meta={
"section_id": f"{section_id}-part-{idx}",
"heading": _text(child_heading) if child_heading else parent_title,
"heading_level": infer_heading_level(child_heading) if child_heading else parent_level,
"parent_heading": parent_title,
"chunk_strategy": "heading_hierarchy_part",
"part_index": idx,
"is_sub_chunk": False,
"retrieval": True,
},
)
)
return parts
def split_by_heading_hierarchy(blocks: list[Block], config: SplitConfig) -> list[ChunkGroup]:
"""
Heading-first splitting:
- Same heading section stays together (body, images, tables, captions).
- Boundaries at same/higher-level headings.
- Oversized sections are split by sub-headings / markers / length (no parent+child pair).
"""
if not blocks:
return []
if not has_meaningful_headings(blocks):
return []
groups: list[ChunkGroup] = []
sections = _split_primary_sections(blocks)
for i, section in enumerate(sections):
heading = _section_heading(section)
sid = _section_key(heading, i)
title = _text(heading) if heading else ""
level = infer_heading_level(heading) if heading else 1
if _rendered_len(section) <= config.max_chunk_size:
groups.append(
ChunkGroup(
blocks=section,
meta={
"section_id": sid,
"heading": title,
"heading_level": level,
"is_sub_chunk": False,
"chunk_strategy": "heading_hierarchy",
"retrieval": True,
"order_range": [
section[0].meta.get("order_index"),
section[-1].meta.get("order_index"),
],
},
)
)
continue
groups.extend(_split_oversized_section(section, config, sid))
return groups
def chunk_groups_to_block_groups(groups: list[ChunkGroup]) -> tuple[list[list[Block]], list[dict]]:
"""Convert ChunkGroups to block groups + per-chunk meta for renderer."""
block_groups: list[list[Block]] = []
metas: list[dict] = []
for group in groups:
block_groups.append(group.blocks)
metas.append(group.meta)
return block_groups, metas
def assign_parent_chunk_ids(chunks: list) -> list:
"""Resolve parent_section_id -> parent_chunk_id (chunk index)."""
section_index: dict[str, int] = {}
for i, chunk in enumerate(chunks):
sid = chunk.meta.get("section_id")
if sid and chunk.meta.get("is_section_parent"):
section_index[sid] = i
updated = []
for chunk in chunks:
meta = dict(chunk.meta)
parent_sid = meta.get("parent_section_id")
if parent_sid and parent_sid in section_index:
meta["parent_chunk_id"] = section_index[parent_sid]
updated.append(chunk.model_copy(update={"meta": meta}))
return updated