Port docs-mcp-template upgrades (eval, citations, chunking, prefixes) (#16)
Co-authored-by: claude <[email protected]>
This commit was merged in pull request #16.
This commit is contained in:
+168
-47
@@ -1,24 +1,12 @@
|
||||
"""Markdown chunker — paragraph-aware, ~400-600 token target.
|
||||
"""Markdown chunker — heading-recursive, ~400-600 token target.
|
||||
|
||||
Adjust the chunking strategy per product if your page format differs
|
||||
significantly from prose. The output shape (id, text, metadata) is
|
||||
fixed by the downstream Chroma + BM25 indexing in rag/index.py — don't
|
||||
change that.
|
||||
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.
|
||||
|
||||
The key knob you'll tune per product is chunk-0. Dense retrieval lands
|
||||
on chunk 0 first for most queries. Make it a synthetic chunk built
|
||||
from:
|
||||
|
||||
- the page title (as natural-language H1)
|
||||
- a 1-sentence task description (you'll have to generate this — for
|
||||
pages that already have a "## Overview" or "## Introduction" the
|
||||
first sentence usually works)
|
||||
- a keyword bag of important terms (filenames, API names, error
|
||||
codes — the rare technical tokens that BM25 lights up on)
|
||||
|
||||
Without a rich chunk 0, dense retrieval gets dominated by the much
|
||||
larger prose body, and short pages (script examples, reference cards)
|
||||
get buried.
|
||||
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
|
||||
|
||||
@@ -26,17 +14,12 @@ import re
|
||||
from typing import Iterator
|
||||
|
||||
|
||||
# Approximate token estimate from char count. Tunable — set per
|
||||
# embedder if the default 4 chars/token is wrong.
|
||||
CHARS_PER_TOKEN = 4
|
||||
TARGET_TOKENS = 500
|
||||
TARGET_CHARS = TARGET_TOKENS * CHARS_PER_TOKEN
|
||||
# Hard cap: nomic-embed-text's context is 2048 tokens. Anything larger
|
||||
# 400s the entire embed batch. 6000 chars works for prose but markdown
|
||||
# tables with lots of `|` separators tokenize ~1.4× denser; a 5839-char
|
||||
# table chunk from the HVM qualification matrix tokenized past 2048 and
|
||||
# crashed the rebuild. 4000 chars stays under 2048 tokens even for
|
||||
# dense table content while leaving headroom for the query side.
|
||||
# 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
|
||||
|
||||
|
||||
@@ -97,6 +80,158 @@ def split_paragraphs(md: str) -> list[str]:
|
||||
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,
|
||||
@@ -112,7 +247,6 @@ def chunks_from_page(
|
||||
if not paragraphs:
|
||||
return
|
||||
|
||||
# ----- Chunk 0: synthetic anchor for dense retrieval ---------
|
||||
title = metadata.get("title") or page_id
|
||||
first_para = next((p for p in paragraphs if not p.startswith("#")), "")
|
||||
chunk0_body = (
|
||||
@@ -127,28 +261,15 @@ def chunks_from_page(
|
||||
"metadata": {**metadata, "ordinal": 0},
|
||||
}
|
||||
|
||||
# ----- Body chunks: pack paragraphs up to TARGET_CHARS -------
|
||||
ordinal = 1
|
||||
|
||||
def emit(buf: list[str]) -> Iterator[dict]:
|
||||
nonlocal ordinal
|
||||
merged = "\n\n".join(buf)
|
||||
for piece in _hard_split(merged):
|
||||
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": {**metadata, "ordinal": ordinal},
|
||||
"metadata": meta,
|
||||
}
|
||||
ordinal += 1
|
||||
|
||||
buf: list[str] = []
|
||||
buf_chars = 0
|
||||
for p in paragraphs:
|
||||
if buf_chars + len(p) > TARGET_CHARS and buf:
|
||||
yield from emit(buf)
|
||||
buf = []
|
||||
buf_chars = 0
|
||||
buf.append(p)
|
||||
buf_chars += len(p)
|
||||
if buf:
|
||||
yield from emit(buf)
|
||||
|
||||
@@ -41,6 +41,21 @@ def _resolve_urls() -> list[str]:
|
||||
OLLAMA_URLS = _resolve_urls()
|
||||
EMBED_MODEL = os.environ.get("EMBED_MODEL", "nomic-embed-text")
|
||||
EMBED_DIM = int(os.environ.get("EMBED_DIM", "768"))
|
||||
# nomic asymmetric prefixes — embed time only. Empty string disables.
|
||||
EMBED_QUERY_PREFIX = os.environ.get("EMBED_QUERY_PREFIX", "search_query: ")
|
||||
EMBED_DOC_PREFIX = os.environ.get("EMBED_DOC_PREFIX", "search_document: ")
|
||||
|
||||
|
||||
def apply_prefix(text: str, prefix: str) -> str:
|
||||
if not prefix:
|
||||
return text
|
||||
return prefix + text
|
||||
|
||||
|
||||
def embed_texts(texts: list[str], *, prefix: str, ef: Any | None = None) -> list[list[float]]:
|
||||
"""Embed texts after applying prefix. Callers store the original strings."""
|
||||
fn = ef if ef is not None else embedding_function()
|
||||
return list(fn([apply_prefix(t, prefix) for t in texts]))
|
||||
|
||||
|
||||
class OllamaEmbeddings(EmbeddingFunction):
|
||||
|
||||
+5
-2
@@ -18,7 +18,7 @@ import chromadb
|
||||
from chromadb.config import Settings
|
||||
|
||||
from .chunk import chunks_from_page
|
||||
from .embeddings import embedding_function
|
||||
from .embeddings import EMBED_DOC_PREFIX, embed_texts, embedding_function
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
|
||||
@@ -78,9 +78,12 @@ def upsert_to_chroma(records: list[dict]) -> int:
|
||||
total = 0
|
||||
for i in range(0, len(records), BATCH):
|
||||
chunk = records[i:i + BATCH]
|
||||
texts = [r["text"] for r in chunk]
|
||||
vectors = embed_texts(texts, prefix=EMBED_DOC_PREFIX, ef=embedding_function())
|
||||
col.upsert(
|
||||
ids=[r["id"] for r in chunk],
|
||||
documents=[r["text"] for r in chunk],
|
||||
documents=texts,
|
||||
embeddings=vectors,
|
||||
metadatas=[r["metadata"] for r in chunk],
|
||||
)
|
||||
total += len(chunk)
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
"""Heading-recursive chunker tests. No corpus, no Chroma."""
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
import rag.chunk as chunk
|
||||
|
||||
|
||||
META = {"bundle_id": "Admin.10.0", "title": "Admin Guide"}
|
||||
|
||||
|
||||
def _bodies(text: str) -> list[str]:
|
||||
return [c["text"] for c in chunk.chunks_from_page(text, "Page", META)]
|
||||
|
||||
|
||||
class ChunkTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self._saved = chunk.TARGET_CHARS
|
||||
|
||||
def tearDown(self) -> None:
|
||||
chunk.TARGET_CHARS = self._saved
|
||||
|
||||
def test_chunk0_always_has_title_on_short_page(self) -> None:
|
||||
page = "# Install\n\nJust one short paragraph about installing."
|
||||
chunks = list(chunk.chunks_from_page(page, "Install", META))
|
||||
self.assertGreaterEqual(len(chunks), 1)
|
||||
self.assertEqual(chunks[0]["metadata"]["ordinal"], 0)
|
||||
self.assertIn("# Admin Guide", chunks[0]["text"])
|
||||
|
||||
def test_sibling_h2_sections_do_not_mix(self) -> None:
|
||||
chunk.TARGET_CHARS = 80
|
||||
page = (
|
||||
"# A\n\n"
|
||||
"intro text for A\n\n"
|
||||
"## A.1\n\n"
|
||||
"alpha content lives here only\n\n"
|
||||
"## A.2\n\n"
|
||||
"beta content lives here only\n"
|
||||
)
|
||||
bodies = _bodies(page)
|
||||
self.assertTrue(any("alpha content" in b and "beta content" not in b for b in bodies[1:]),
|
||||
bodies)
|
||||
self.assertTrue(any("beta content" in b and "alpha content" not in b for b in bodies[1:]),
|
||||
bodies)
|
||||
|
||||
def test_giant_section_splits_on_paragraphs(self) -> None:
|
||||
chunk.TARGET_CHARS = 40
|
||||
paras = [f"Paragraph number {i} with enough words." for i in range(8)]
|
||||
page = "# Giant\n\n" + "\n\n".join(paras)
|
||||
bodies = _bodies(page)
|
||||
# skip chunk 0
|
||||
for b in bodies[1:]:
|
||||
# may exceed by at most one paragraph (the one that filled the buf)
|
||||
self.assertLessEqual(len(b), chunk.TARGET_CHARS + len(paras[0]) + 2, b)
|
||||
|
||||
def test_fenced_code_never_sliced(self) -> None:
|
||||
chunk.TARGET_CHARS = 30
|
||||
fence = "```\n" + ("x" * 80) + "\n```"
|
||||
page = "# Code\n\nBefore.\n\n" + fence + "\n\nAfter."
|
||||
bodies = _bodies(page)
|
||||
joined = "\n".join(bodies)
|
||||
self.assertIn(fence, joined)
|
||||
# no body chunk should contain a half-fence
|
||||
for b in bodies:
|
||||
if "```" in b:
|
||||
self.assertTrue(b.strip().startswith("```") or "```\n" in b)
|
||||
self.assertGreaterEqual(b.count("```"), 2, b)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Prefix helpers — no Ollama, no Chroma."""
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from rag.embeddings import apply_prefix, embed_texts
|
||||
|
||||
|
||||
class FakeEmbedder:
|
||||
def __init__(self) -> None:
|
||||
self.seen: list[str] = []
|
||||
|
||||
def __call__(self, texts: list[str]) -> list[list[float]]:
|
||||
self.seen.extend(texts)
|
||||
return [[0.0, 1.0] for _ in texts]
|
||||
|
||||
|
||||
class PrefixTests(unittest.TestCase):
|
||||
def test_index_sends_doc_prefix_keeps_raw_text(self) -> None:
|
||||
fake = FakeEmbedder()
|
||||
raw = ["hello"]
|
||||
vecs = embed_texts(raw, prefix="search_document: ", ef=fake)
|
||||
self.assertEqual(fake.seen, ["search_document: hello"])
|
||||
self.assertEqual(raw, ["hello"]) # caller storage unchanged
|
||||
self.assertEqual(len(vecs), 1)
|
||||
|
||||
def test_query_sends_query_prefix(self) -> None:
|
||||
fake = FakeEmbedder()
|
||||
embed_texts(["q"], prefix="search_query: ", ef=fake)
|
||||
self.assertEqual(fake.seen, ["search_query: q"])
|
||||
self.assertNotIn("q", fake.seen)
|
||||
|
||||
def test_empty_prefix_is_raw(self) -> None:
|
||||
self.assertEqual(apply_prefix("hello", ""), "hello")
|
||||
fake = FakeEmbedder()
|
||||
embed_texts(["hello"], prefix="", ef=fake)
|
||||
self.assertEqual(fake.seen, ["hello"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user