Files
RAG-CUT/backend/rag_cut/splitters/heading_splitter.py
T

399 lines
13 KiB
Python
Raw Normal View History

2026-07-16 11:12:17 +08:00
"""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.上升三角形態".
# The title starts with a letter/CJK character so "1.1.1" cannot backtrack
# into number="1" + title="1.1".
2026-07-16 11:12:17 +08:00
NUMBERED_HEADING_RE = re.compile(
r"^\s*(\d+(?:\.\d+)*)(?:[..、::]\s*|\s+)"
r"([A-Za-z\u4e00-\u9fff][^\n]{1,120})\s*$"
)
PURE_NUMBERED_HEADING_RE = re.compile(r"^\s*(\d+(?:\.\d+){1,5})\.?\s*$")
TIME_LIKE_RE = re.compile(r"^\s*\d{1,2}:\d{2}(?::\d{2})?\s*$")
NUMBERED_LIST_ITEM_RE = re.compile(r"^\s*\d+[.).、]\s+(?P<body>.+)$")
NUMBERED_OPTION_SENTENCE_RE = re.compile(
r"^\s*\d+(?:\.\d+)+\.?\s+.+\bthis\s+(?:option|function|feature)\s+"
r"(?:enables|allows)\b",
re.I,
)
INSTRUCTION_START_RE = re.compile(
r"^(?:select|click|choose|enter|input|open|close|press|perform|to\s+|on\s+|"
r"the\s+user|users?\s+|next\s+|then\s+|for\s+ease|option/tool|"
r"用户|点击|輸入|输入|選擇|选择|填写|當|当|在|首先|然后|然後|配置|系统|系統|"
r"若|如果|注|详细|詳細|列表|机构|機構|添加|删除|刪除)",
re.I,
2026-07-16 11:12:17 +08:00
)
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:
stripped = text.strip()
match = PURE_NUMBERED_HEADING_RE.match(stripped)
if match:
return match.group(1).count(".") + 1
match = NUMBERED_HEADING_RE.match(stripped)
2026-07-16 11:12:17 +08:00
if not match:
return None
return match.group(1).count(".") + 1
def looks_like_numbered_instruction(text: str) -> bool:
"""True for numbered procedure/list sentences, not outline headings."""
stripped = text.strip()
if NUMBERED_OPTION_SENTENCE_RE.match(stripped):
return True
match = NUMBERED_LIST_ITEM_RE.match(stripped)
if not match:
return False
body = match.group("body").strip()
if INSTRUCTION_START_RE.match(body):
return True
return len(body) >= 45 and bool(re.search(r"[,.,。;;:]", body))
def looks_like_false_heading_text(text: str) -> bool:
stripped = text.strip()
if TIME_LIKE_RE.match(stripped) or looks_like_numbered_instruction(stripped):
return True
if numbered_heading_level(stripped):
return False
return len(stripped) >= 40 and bool(re.search(r"[.!?。!?;;]\s*$", stripped))
2026-07-16 11:12:17 +08:00
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 looks_like_false_heading_text(text):
return False
2026-07-16 11:12:17 +08:00
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)
false_heading = looks_like_false_heading_text(text)
if block.type == BlockType.HEADING and (
len(text) > MAX_HEADING_CHARS or false_heading
):
2026-07-16 11:12:17 +08:00
# Parser sometimes merges title+body then marks the blob as heading.
return Block(type=BlockType.PARAGRAPH, text=text, meta=dict(block.meta))
if false_heading:
return block
2026-07-16 11:12:17 +08:00
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) or PURE_NUMBERED_HEADING_RE.match(text)
):
2026-07-16 11:12:17 +08:00
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