276 lines
8.5 KiB
Python
276 lines
8.5 KiB
Python
"""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 dense table chunks.
|
||
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
|