Files
hvm-docs/rag/chunk.py
T

276 lines
8.5 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.
"""Markdown chunker — heading-recursive, ~400-600 token target.
Chunk by semantic section (ATX headings), not raw page/length. A
synthetic chunk 0 (title + first paragraph + optional keyword bag) is
always emitted first — dense retrieval lands on it. Do not drop it.
No chonkie / langchain dependency. The output shape (id, text, metadata)
is fixed by rag/index.py — don't change that. `heading_path` is optional
metadata (e.g. "Install > Linux") and is not required for indexing.
"""
from __future__ import annotations
import re
from typing import Iterator
CHARS_PER_TOKEN = 4
TARGET_TOKENS = 500
TARGET_CHARS = TARGET_TOKENS * CHARS_PER_TOKEN
# nomic-embed-text context is 2048 tokens. Markdown tables with lots of
# `|` tokenize ~1.4× denser than prose; 4000 chars stays under 2048 even
# for qualification-matrix chunks (a 5839-char table crashed a rebuild).
MAX_CHARS = 4000
def _hard_split(text: str) -> list[str]:
"""Split an oversized block on line boundaries into MAX_CHARS pieces."""
if len(text) <= MAX_CHARS:
return [text]
out: list[str] = []
buf: list[str] = []
buf_chars = 0
for line in text.splitlines(keepends=True):
if buf_chars + len(line) > MAX_CHARS and buf:
out.append("".join(buf).rstrip())
buf, buf_chars = [], 0
buf.append(line)
buf_chars += len(line)
if buf:
out.append("".join(buf).rstrip())
return out
def estimate_tokens(text: str) -> int:
return max(1, len(text) // CHARS_PER_TOKEN)
def split_paragraphs(md: str) -> list[str]:
"""Split markdown into paragraph-ish blocks.
Keeps fenced code blocks together (don't slice through ```).
Headings start new paragraphs.
"""
blocks: list[str] = []
current: list[str] = []
in_fence = False
for line in md.splitlines(keepends=True):
stripped = line.strip()
if stripped.startswith("```"):
in_fence = not in_fence
current.append(line)
continue
if in_fence:
current.append(line)
continue
if stripped.startswith("#"):
if current:
blocks.append("".join(current).strip())
current = []
current.append(line)
continue
if not stripped and current and not "".join(current).strip().endswith("\n\n"):
current.append(line)
blocks.append("".join(current).strip())
current = []
continue
current.append(line)
if current:
blocks.append("".join(current).strip())
return [b for b in blocks if b]
def _heading_level(block: str) -> int:
first = block.lstrip().split("\n", 1)[0].strip()
m = re.match(r"^(#{1,6})\s+\S", first)
return len(m.group(1)) if m else 0
def _heading_title(block: str) -> str:
first = block.split("\n", 1)[0].strip()
return re.sub(r"^#{1,6}\s+", "", first).strip()
def _is_fence(block: str) -> bool:
return block.lstrip().startswith("```")
def _joined_len(blocks: list[str]) -> int:
if not blocks:
return 0
return sum(len(b) for b in blocks) + 2 * (len(blocks) - 1)
def _split_len(text: str) -> list[str]:
size = TARGET_CHARS
return [text[i:i + size] for i in range(0, len(text), size)] or [text]
def _hard_wrap(text: str) -> list[str]:
"""Last-resort split. Never used on fenced code blocks."""
if len(text) <= TARGET_CHARS:
return [text]
parts = re.split(r"\n\s*\n", text)
if len(parts) <= 1:
return _split_len(text)
packed: list[str] = []
buf: list[str] = []
n = 0
for part in parts:
if len(part) > TARGET_CHARS:
if buf:
packed.append("\n\n".join(buf))
buf, n = [], 0
packed.extend(_split_len(part))
continue
if n + len(part) > TARGET_CHARS and buf:
packed.append("\n\n".join(buf))
buf, n = [], 0
buf.append(part)
n += len(part)
if buf:
packed.append("\n\n".join(buf))
return packed
def _pack_paragraphs(blocks: list[str], path: str) -> list[tuple[str, str]]:
out: list[tuple[str, str]] = []
buf: list[str] = []
buf_chars = 0
def flush() -> None:
nonlocal buf, buf_chars
if buf:
out.append(("\n\n".join(buf), path))
buf, buf_chars = [], 0
for p in blocks:
if _is_fence(p) and len(p) > TARGET_CHARS:
flush()
out.append((p, path)) # never slice a fence
continue
if len(p) > TARGET_CHARS:
flush()
out.extend((piece, path) for piece in _hard_wrap(p))
continue
if buf_chars + len(p) > TARGET_CHARS and buf:
flush()
buf.append(p)
buf_chars += len(p)
flush()
return out
def _split_at_level(blocks: list[str], level: int) -> list[tuple[list[str], list[str]]]:
"""Split into (path_suffix, blocks) groups starting at `level` headings.
path_suffix is [] for preamble, [title] for a section headed at `level`.
"""
groups: list[tuple[list[str], list[str]]] = []
preamble: list[str] = []
current_path: list[str] = []
current: list[str] = []
for b in blocks:
lv = _heading_level(b)
if lv == level:
if current:
groups.append((current_path, current))
elif preamble:
groups.append(([], preamble))
preamble = []
current_path = [_heading_title(b)]
current = [b]
elif current:
current.append(b)
else:
preamble.append(b)
if current:
if preamble and not groups:
groups.append(([], preamble))
preamble = []
groups.append((current_path, current))
elif preamble:
groups.append(([], preamble))
return groups
def _pack_recursive(blocks: list[str], path: list[str]) -> list[tuple[str, str]]:
if not blocks:
return []
path_s = " > ".join(path)
if _joined_len(blocks) <= TARGET_CHARS:
return [("\n\n".join(blocks), path_s)]
levels = [_heading_level(b) for b in blocks]
heading_levels = [lv for lv in levels if lv > 0]
if not heading_levels:
return _pack_paragraphs(blocks, path_s)
split_lv = min(heading_levels)
groups = _split_at_level(blocks, split_lv)
if len(groups) <= 1:
# Can't split at this heading level — descend or pack paragraphs.
if levels[0] == split_lv:
title = _heading_title(blocks[0])
rest = blocks[1:]
if not rest:
return _pack_paragraphs(blocks, path_s)
packed = _pack_recursive(rest, path + [title])
if not packed:
return [(blocks[0], " > ".join(path + [title]))]
glued = blocks[0] + "\n\n" + packed[0][0]
if len(glued) <= TARGET_CHARS:
packed[0] = (glued, packed[0][1])
else:
packed.insert(0, (blocks[0], packed[0][1]))
return packed
return _pack_paragraphs(blocks, path_s)
out: list[tuple[str, str]] = []
for suffix, group in groups:
out.extend(_pack_recursive(group, path + suffix))
return out
def chunks_from_page(
text: str,
page_id: str,
metadata: dict,
) -> Iterator[dict]:
"""Yield chunk dicts ready for index.py to upsert.
The synthetic chunk 0 is the per-product customization point. The
default below is a simple title + body-first-paragraph; rewrite
for richer retrieval signal (see module docstring).
"""
paragraphs = split_paragraphs(text)
if not paragraphs:
return
title = metadata.get("title") or page_id
first_para = next((p for p in paragraphs if not p.startswith("#")), "")
chunk0_body = (
f"# {title}\n\n"
f"{first_para[:300]}"
# TODO per product: append a keyword bag here (filenames,
# API names, error codes) for BM25 + dense joint coverage.
)
yield {
"id": f"{metadata['bundle_id']}::{page_id}::0",
"text": chunk0_body,
"metadata": {**metadata, "ordinal": 0},
}
ordinal = 1
for body, heading_path in _pack_recursive(paragraphs, []):
for piece in _hard_split(body):
meta = {**metadata, "ordinal": ordinal}
if heading_path:
meta["heading_path"] = heading_path
yield {
"id": f"{metadata['bundle_id']}::{page_id}::{ordinal}",
"text": piece,
"metadata": meta,
}
ordinal += 1