Port docs-mcp-template upgrades (eval, citations, chunking, prefixes) #13

Merged
claude merged 1 commits from issue-12 into main 2026-09-30 08:33:17 -04:00
13 changed files with 812 additions and 77 deletions
Showing only changes of commit fda572f59e - Show all commits
+43
View File
@@ -0,0 +1,43 @@
"""Markdown formatters for MCP tool output.
No third-party imports — tests can run without mcp/httpx/chromadb.
Citation numbers are per-call (1-based, dense) and stateless.
"""
from __future__ import annotations
def format_search_hit(
n: int,
title: str,
url: str,
text: str,
extra: str = "",
) -> str:
"""One numbered hit. `n` is 1-based in final reranked/fused order."""
head = f"[{n}] **{title}**"
if url:
head += f" — {url}"
parts = [head]
if extra:
parts.append(extra)
body = (text or "").strip()
if body:
parts.append(body)
return "\n".join(parts)
def format_search_hits(
hits: list[tuple[str, str, str, str]],
) -> str:
"""Render hits as `[1] **title** — url` then text.
`hits` is a list of (title, url, text, extra) in display order.
Empty list → empty string (no invented citations).
"""
if not hits:
return ""
blocks = [
format_search_hit(n, title, url, text, extra)
for n, (title, url, text, extra) in enumerate(hits, start=1)
]
return "\n\n".join(blocks) + "\n"
+29 -17
View File
@@ -31,6 +31,7 @@ from mcp.server.mcpserver import MCPServer
from mcp.server.transport_security import TransportSecuritySettings
from pydantic import Field
from .format import format_search_hits
from .usage import TimedCall
log = logging.getLogger(__name__)
@@ -145,6 +146,17 @@ def _collection():
return _CHROMA
def _query_dense(col, query: str, n: int, where: dict | None = None, **extra):
"""Dense query with nomic prefixes applied at embed time only."""
from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts
qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0]
kwargs: dict[str, Any] = {"query_embeddings": [qvec], "n_results": n}
if where:
kwargs["where"] = where
kwargs.update(extra)
return col.query(**kwargs)
def _bm25():
"""Lazy BM25Index handle. None if the FTS5 db isn't built."""
global _BM25
@@ -258,9 +270,10 @@ def search_docs(
"""Search the HPE Morpheus Enterprise (Morpheus) docs corpus.
Returns the top-k most relevant chunks (with full source page URLs)
given a natural-language query. Optional filters narrow the search
to one version, one platform, or one bundle. Use list_versions()
first if you need to discover the available facet values.
given a natural-language query. Hits are numbered [1]… for citation.
Optional filters narrow the search to one version, one platform, or
one bundle. Use list_versions() first if you need to discover the
available facet values.
Call this tool whenever the user asks anything that should be
answerable from the official product documentation — install,
@@ -296,7 +309,7 @@ def search_docs(
if HYBRID_SEARCH and bm is not None:
try:
dense_res = col.query(query_texts=[query], n_results=pool, where=where)
dense_res = _query_dense(col, query, pool, where)
dense_ids = (dense_res.get("ids") or [[]])[0]
bm_hits = bm.query(query, n=pool, where=bm25_where)
bm_ids = [cid for cid, _s in bm_hits]
@@ -326,7 +339,7 @@ def search_docs(
log.warning("BM25 retrieval failed, falling back to dense: %s", e)
if not docs:
res = col.query(query_texts=[query], n_results=k, where=where)
res = _query_dense(col, query, k, where)
docs = (res.get("documents") or [[]])[0]
metas = (res.get("metadatas") or [[]])[0]
dists = (res.get("distances") or [[]])[0]
@@ -342,7 +355,7 @@ def search_docs(
extra = bm.query(query, n=pool_size, where=bm25_where) if bm else []
extra_ids = [cid for cid, _s in extra]
else:
extra_res = col.query(query_texts=[query], n_results=pool_size, where=where)
extra_res = _query_dense(col, query, pool_size, where)
extra_ids = (extra_res.get("ids") or [[]])[0]
if extra_ids:
d2, m2, _ = _enrich_from_chroma(col, extra_ids, None)
@@ -371,22 +384,21 @@ def search_docs(
if not docs:
return f"_No matches for `{query}`._"
out = [f"# {len(docs)} result(s) for `{query}`", ""]
hits: list[tuple[str, str, str, str]] = []
for doc, meta, dist in zip(docs, metas, dists):
bid = meta.get("bundle_id", "")
pid = meta.get("page_id", "")
title = meta.get("title") or pid
ver = meta.get("version") or ""
url = _source_url(bid, pid)
header = f"## {title}"
extra = f"score={1 - dist:.3f}"
if ver:
header += f" _(v{ver})_"
out.append(header)
out.append(f"[{bid}/{pid}]({url}) · score={1 - dist:.3f}")
out.append("")
out.append(doc.strip())
out.append("")
return "\n".join(out)
extra = f"v{ver} " + extra
if pid:
extra += f" page_id: `{pid}`"
hits.append((title, url, doc, extra))
header = f"# {len(docs)} result(s) for `{query}`\n\n"
return header + format_search_hits(hits)
@mcp.tool()
@@ -1073,8 +1085,8 @@ def find_doc_inconsistencies(
return f"Couldn't open Chroma collection: {e}"
where = _build_where(version, platform, bundle_id)
try:
res = col.query(query_texts=[scope_query], n_results=max_pages * 3,
where=where, include=["metadatas"])
res = _query_dense(col, scope_query, max_pages * 3, where,
include=["metadatas"])
except Exception as e:
_call.set(error=f"query: {e}")
return f"Scope query failed: {e}"
+37
View File
@@ -0,0 +1,37 @@
"""Citation markdown renderer — no Chroma, no mcp."""
from __future__ import annotations
import unittest
from docs_mcp.format import format_search_hits
class FormatSearchHitsTests(unittest.TestCase):
def test_two_hits_numbered_with_urls(self) -> None:
out = format_search_hits([
("Install on Linux", "https://docs.example.com/install", "chunk one", ""),
("HA setup", "https://docs.example.com/ha", "chunk two", ""),
])
self.assertIn("[1] **Install on Linux** — https://docs.example.com/install", out)
self.assertIn("[2] **HA setup** — https://docs.example.com/ha", out)
self.assertIn("chunk one", out)
self.assertIn("chunk two", out)
self.assertNotIn("[3]", out)
# [1] before [2]
self.assertLess(out.index("[1]"), out.index("[2]"))
# each number is followed by its URL on the same logical hit
first = out.split("[2]")[0]
self.assertIn("https://docs.example.com/install", first)
self.assertNotIn("https://docs.example.com/ha", first)
def test_empty_hits_invents_nothing(self) -> None:
self.assertEqual(format_search_hits([]), "")
def test_missing_url_omits_emdash(self) -> None:
out = format_search_hits([("Title only", "", "body", "")])
self.assertIn("[1] **Title only**", out)
self.assertNotIn("—", out)
if __name__ == "__main__":
unittest.main()
+156
View File
@@ -0,0 +1,156 @@
"""Paired permutation test between two eval JSONL sidecars.
Compares per-query scores from two `eval.run_eval` sidecar files so
"P@1 went 0.88 → 0.91 on 25 queries" is not treated as a win.
python -m eval.pvalue \\
--a eval/results/baseline.jsonl \\
--b eval/results/new.jsonl \\
--metric rr
Exit 0 even when the difference is not significant — this is a report,
not a gate. No third-party deps; `random.Random(seed)` is enough.
"""
from __future__ import annotations
import argparse
import json
import random
from pathlib import Path
def load_sidecar(path: Path) -> list[dict]:
rows: list[dict] = []
with open(path) as fh:
for line in fh:
line = line.strip()
if line:
rows.append(json.loads(line))
return rows
def paired_permutation(
a: list[float],
b: list[float],
n_resamples: int = 10000,
seed: int = 0,
) -> dict:
"""Two-sided paired permutation test on per-query scores.
Null: each pair is exchangeable (randomly flipping the sign of
A_i - B_i). p_value is the fraction of permutations whose
|mean diff| is at least as large as the observed |mean(A-B)|.
"""
if len(a) != len(b):
raise ValueError(f"paired lengths differ: {len(a)} vs {len(b)}")
if not a:
raise ValueError("no paired queries to compare")
diffs = [x - y for x, y in zip(a, b)]
n = len(diffs)
observed = sum(diffs) / n
abs_obs = abs(observed)
rng = random.Random(seed)
extreme = 0
for _ in range(n_resamples):
total = 0.0
for d in diffs:
total += d if rng.random() < 0.5 else -d
if abs(total / n) >= abs_obs - 1e-15:
extreme += 1
p_value = extreme / n_resamples
return {
"A_mean": sum(a) / n,
"B_mean": sum(b) / n,
"Diff(A-B)": observed,
"p_value": p_value,
"significant": p_value < 0.05,
"n": n,
"n_resamples": n_resamples,
}
def _index(rows: list[dict], metric: str) -> dict[tuple[str, str], float]:
"""Map (retriever, query) -> score."""
out: dict[tuple[str, str], float] = {}
for row in rows:
retriever = str(row.get("retriever") or "")
query = str(row.get("query") or "")
if metric == "p_at_1":
score = float(row.get("p_at_1") or 0)
else:
score = float(row.get("rr") or 0)
out[(retriever, query)] = score
return out
def compare(
rows_a: list[dict],
rows_b: list[dict],
metric: str = "rr",
retriever: str | None = None,
n_resamples: int = 10000,
seed: int = 0,
) -> list[dict]:
"""Join on (retriever, query). One result dict per shared retriever."""
ia, ib = _index(rows_a, metric), _index(rows_b, metric)
retrievers = sorted({r for r, _ in ia} & {r for r, _ in ib})
if retriever:
retrievers = [r for r in retrievers if r == retriever]
if not retrievers:
raise ValueError(f"retriever {retriever!r} not in both sidecars")
reports = []
for name in retrievers:
queries = sorted({q for r, q in ia if r == name} & {q for r, q in ib if r == name})
if not queries:
continue
a_scores = [ia[(name, q)] for q in queries]
b_scores = [ib[(name, q)] for q in queries]
report = paired_permutation(a_scores, b_scores, n_resamples=n_resamples, seed=seed)
report["retriever"] = name
report["metric"] = metric
reports.append(report)
if not reports:
raise ValueError("no overlapping (retriever, query) pairs")
return reports
def render(reports: list[dict]) -> str:
lines = ["# Permutation test", ""]
for r in reports:
sig = "yes" if r["significant"] else "no"
lines += [
f"## `{r['retriever']}` ({r['metric']}, n={r['n']})",
"",
f"- A_mean: `{r['A_mean']:.4f}`",
f"- B_mean: `{r['B_mean']:.4f}`",
f"- Diff(A-B): `{r['Diff(A-B)']:.4f}`",
f"- p_value: `{r['p_value']:.4f}` ({r['n_resamples']} resamples)",
f"- significant (p < 0.05): **{sig}**",
"",
]
return "\n".join(lines)
def main() -> int:
p = argparse.ArgumentParser(description="Paired permutation test on two eval JSONL sidecars.")
p.add_argument("--a", type=Path, required=True, help="sidecar JSONL (system A)")
p.add_argument("--b", type=Path, required=True, help="sidecar JSONL (system B)")
p.add_argument("--metric", choices=("rr", "p_at_1"), default="rr")
p.add_argument("--retriever", default=None, help="restrict to one retriever name")
p.add_argument("--n-resamples", type=int, default=10000)
p.add_argument("--seed", type=int, default=0)
args = p.parse_args()
reports = compare(
load_sidecar(args.a),
load_sidecar(args.b),
metric=args.metric,
retriever=args.retriever,
n_resamples=args.n_resamples,
seed=args.seed,
)
print(render(reports), end="")
return 0
if __name__ == "__main__":
raise SystemExit(main())
+3 -1
View File
@@ -51,7 +51,9 @@ class DenseRetriever:
self.pool = pool
def retrieve(self, query: str, k: int = 10) -> list[tuple[str, str]]:
res = self.col.query(query_texts=[query], n_results=self.pool)
from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts
qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0]
res = self.col.query(query_embeddings=[qvec], n_results=self.pool)
ids = (res.get("ids") or [[]])[0]
return _collapse_to_pages(ids, k)
+38 -8
View File
@@ -34,6 +34,12 @@ def load_queries(path: Path) -> list[dict]:
return [json.loads(line) for line in fh if line.strip()]
def p_at_1(retrieved: list[tuple[str, str]], expected: list[tuple[str, str]]) -> float:
if not retrieved or not expected:
return 0.0
return 1.0 if retrieved[0] in set(expected) else 0.0
def reciprocal_rank(retrieved: list[tuple[str, str]], expected: list[tuple[str, str]]) -> float:
expected_set = set(expected)
for i, page in enumerate(retrieved, start=1):
@@ -65,8 +71,12 @@ def main() -> int:
p = argparse.ArgumentParser()
p.add_argument("--queries", type=Path, default=Path("eval/queries.jsonl"))
p.add_argument("--k", type=int, default=5)
p.add_argument("--ks", default="1,5,10,20", help="comma-separated k-curve")
p.add_argument("--output", type=Path, default=Path("eval/results/baseline.md"))
p.add_argument("--compare", type=Path, default=None)
args = p.parse_args()
ks = sorted({int(x) for x in args.ks.split(",") if x.strip()}) or [args.k]
max_k = max(ks + [args.k])
if not args.queries.exists():
print(f"queries file not found: {args.queries}")
@@ -83,7 +93,7 @@ def main() -> int:
from rag.bm25 import BM25Index
from eval.retrievers import DenseRetriever, BM25Retriever, HybridRetriever
product = os.environ.get("PRODUCT_NAME", "hvm")
product = os.environ.get("PRODUCT_NAME", "morpheus")
repo_root = Path(__file__).resolve().parent.parent
client = chromadb.PersistentClient(path=str(repo_root / "chroma"),
settings=Settings(anonymized_telemetry=False))
@@ -109,34 +119,39 @@ def main() -> int:
rows: dict[str, dict[str, float]] = {}
per_query: list[dict] = []
for r in retrievers:
mrr_sum = recall_sum = ndcg_sum = 0.0
mrr_sum = recall_sum = ndcg_sum = p1_sum = 0.0
elapsed_sum = 0.0
for q in queries:
expected = [(e["bundle_id"], e["page_id"]) for e in q["expected"]]
t0 = time.time()
retrieved = r.retrieve(q["query"], k=max(args.k, 10))
retrieved = r.retrieve(q["query"], k=max(max_k, 10))
elapsed = time.time() - t0
mrr = reciprocal_rank(retrieved, expected)
p1 = p_at_1(retrieved, expected)
recall = recall_at_k(retrieved, expected, args.k)
ndcg = ndcg_at_k(retrieved, expected, args.k)
mrr_sum += mrr
p1_sum += p1
recall_sum += recall
ndcg_sum += ndcg
elapsed_sum += elapsed
per_query.append({
"retriever": r.name, "query": q["query"],
"mrr": mrr, "recall@k": recall, "ndcg@k": ndcg,
"mrr": mrr, "p_at_1": int(p1), "recall@k": recall, "ndcg@k": ndcg,
"top1": list(retrieved[0]) if retrieved else None,
"ranked": [list(p) for p in retrieved],
"elapsed_s": round(elapsed, 3),
})
n = len(queries)
rows[r.name] = {
"P@1": p1_sum / n,
"MRR": mrr_sum / n,
f"Recall@{args.k}": recall_sum / n,
f"nDCG@{args.k}": ndcg_sum / n,
"avg_latency_s": elapsed_sum / n,
}
print(f" {r.name}: MRR={rows[r.name]['MRR']:.3f} "
print(f" {r.name}: P@1={rows[r.name]['P@1']:.3f} "
f"MRR={rows[r.name]['MRR']:.3f} "
f"Recall@{args.k}={rows[r.name][f'Recall@{args.k}']:.3f} "
f"nDCG@{args.k}={rows[r.name][f'nDCG@{args.k}']:.3f} "
f"avg={rows[r.name]['avg_latency_s']*1000:.0f}ms")
@@ -144,10 +159,10 @@ def main() -> int:
args.output.parent.mkdir(parents=True, exist_ok=True)
md = [f"# Retrieval eval — k={args.k}", "",
f"_{len(queries)} hand-curated queries, generated {time.strftime('%Y-%m-%d %H:%M:%S')}_", "",
"| Retriever | MRR | Recall@{k} | nDCG@{k} | avg latency |".replace("{k}", str(args.k)),
"| --- | ---: | ---: | ---: | ---: |"]
"| Retriever | P@1 | MRR | Recall@{k} | nDCG@{k} | avg latency |".replace("{k}", str(args.k)),
"| --- | ---: | ---: | ---: | ---: | ---: |"]
for name, m in rows.items():
md.append(f"| `{name}` | {m['MRR']:.3f} | {m[f'Recall@{args.k}']:.3f} "
md.append(f"| `{name}` | {m['P@1']:.3f} | {m['MRR']:.3f} | {m[f'Recall@{args.k}']:.3f} "
f"| {m[f'nDCG@{args.k}']:.3f} | {m['avg_latency_s']*1000:.0f}ms |")
md += ["", "## Per-query results", "",
"| Retriever | Query | MRR | top-1 |", "| --- | --- | ---: | --- |"]
@@ -155,7 +170,22 @@ def main() -> int:
top1 = f"`{r['top1'][0]}/{r['top1'][1][:24]}...`" if r["top1"] else "—"
md.append(f"| `{r['retriever']}` | {r['query'][:60]} | {r['mrr']:.3f} | {top1} |")
args.output.write_text("\n".join(md) + "\n")
sidecar = args.output.with_suffix(".jsonl")
with open(sidecar, "w") as fh:
for r in per_query:
fh.write(json.dumps({
"query": r["query"],
"retriever": r["retriever"],
"ranked": r.get("ranked") or [],
"rr": r["mrr"],
"p_at_1": r["p_at_1"],
}) + "\n")
print(f"wrote {args.output}")
print(f"wrote {sidecar}")
if args.compare:
from eval.pvalue import compare, load_sidecar, render
print()
print(render(compare(load_sidecar(sidecar), load_sidecar(args.compare))), end="")
return 0
+84
View File
@@ -0,0 +1,84 @@
"""Stdlib tests for eval metrics + the permutation test.
Must not open Chroma — the template has no corpus. Run with:
python -m unittest eval.test_metrics
"""
from __future__ import annotations
import math
import unittest
from eval.pvalue import compare, paired_permutation
from eval.run_eval import ndcg_at_k, p_at_1, recall_at_k, reciprocal_rank
A, B, C, X, Y = ("b", "a"), ("b", "b"), ("b", "c"), ("b", "x"), ("b", "y")
class MetricTests(unittest.TestCase):
def test_reciprocal_rank(self) -> None:
# Q1: expected at rank 1
self.assertEqual(reciprocal_rank([A, B, C], [A]), 1.0)
# Q2: expected at rank 2
self.assertEqual(reciprocal_rank([B, A], [A]), 0.5)
# Q3: miss
self.assertEqual(reciprocal_rank([X, Y], [A]), 0.0)
def test_p_at_1(self) -> None:
self.assertEqual(p_at_1([A, B], [A]), 1.0)
self.assertEqual(p_at_1([B, A], [A]), 0.0)
self.assertEqual(p_at_1([], [A]), 0.0)
self.assertEqual(p_at_1([A], []), 0.0)
def test_recall_at_k(self) -> None:
self.assertEqual(recall_at_k([A, B, C], [A], 1), 1.0)
self.assertEqual(recall_at_k([B, A], [A], 1), 0.0)
self.assertEqual(recall_at_k([B, A], [A], 2), 1.0)
self.assertEqual(recall_at_k([X, Y], [A], 5), 0.0)
self.assertEqual(recall_at_k([A, B], [A, C], 1), 0.5)
def test_ndcg_at_k(self) -> None:
self.assertEqual(ndcg_at_k([A], [A], 1), 1.0)
# expected at rank 2: dcg = 1/log2(3), idcg = 1
self.assertAlmostEqual(
ndcg_at_k([B, A], [A], 2),
(1.0 / math.log2(3)) / 1.0,
)
self.assertEqual(ndcg_at_k([X, Y], [A], 5), 0.0)
class PermutationTests(unittest.TestCase):
def test_identical_lists_not_significant(self) -> None:
scores = [0.5, 1.0, 0.0, 1.0, 0.5]
report = paired_permutation(scores, list(scores), n_resamples=2000, seed=0)
self.assertEqual(report["Diff(A-B)"], 0.0)
self.assertEqual(report["p_value"], 1.0)
self.assertFalse(report["significant"])
def test_large_paired_difference_is_significant(self) -> None:
a = [1.0] * 20
b = [0.0] * 20
report = paired_permutation(a, b, n_resamples=5000, seed=0)
self.assertGreater(report["Diff(A-B)"], 0.9)
self.assertLess(report["p_value"], 0.05)
self.assertTrue(report["significant"])
def test_compare_joins_on_retriever_and_query(self) -> None:
rows_a = (
[{"retriever": "dense", "query": f"q{i}", "rr": 1.0, "p_at_1": 1} for i in range(20)]
+ [{"retriever": "bm25", "query": "q0", "rr": 0.0, "p_at_1": 0}]
)
rows_b = (
[{"retriever": "dense", "query": f"q{i}", "rr": 0.0, "p_at_1": 0} for i in range(20)]
+ [{"retriever": "bm25", "query": "q0", "rr": 0.0, "p_at_1": 0}]
)
reports = compare(rows_a, rows_b, metric="rr", n_resamples=2000, seed=0)
by_name = {r["retriever"]: r for r in reports}
self.assertIn("dense", by_name)
self.assertTrue(by_name["dense"]["significant"])
self.assertFalse(by_name["bm25"]["significant"])
if __name__ == "__main__":
unittest.main()
+114
View File
@@ -0,0 +1,114 @@
"""Page-level miss dump using this clone's Dense/BM25 retrievers.
python -m eval.trace --queries eval/queries.jsonl
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from eval.run_eval import load_queries, p_at_1
def classify_top1_source(top1, dense_pages, bm25_pages) -> str:
if top1 is None:
return "neither"
in_d, in_b = top1 in dense_pages, top1 in bm25_pages
if in_d and in_b:
return "both"
if in_d:
return "dense_only"
if in_b:
return "bm25_only"
return "neither"
def first_ranks(pages: list[tuple[str, str]]) -> dict[str, int]:
out: dict[str, int] = {}
for i, (bid, pid) in enumerate(pages, start=1):
key = f"{bid}/{pid}"
if key not in out:
out[key] = i
return out
def render_misses(rows: list[dict]) -> str:
misses = [r for r in rows if not r.get("hit")]
if not misses:
return "# Eval misses\n\n_(none)_\n"
lines = [f"# Eval misses ({len(misses)})", ""]
for row in misses:
lines += [
f"## {row['query']}",
"",
f"- expected: `{row['expected']}`",
f"- top-5: `{row['ranked_pages'][:5]}`",
f"- top1_source: `{row['top1_source']}`",
"",
]
return "\n".join(lines)
def main() -> int:
p = argparse.ArgumentParser()
p.add_argument("--queries", type=Path, default=Path("eval/queries.jsonl"))
p.add_argument("--trace-out", type=Path, default=Path("eval/results/trace.jsonl"))
p.add_argument("--misses-out", type=Path, default=Path("eval/results/misses.md"))
args = p.parse_args()
if not args.queries.exists():
print(f"queries file not found: {args.queries}")
return 1
try:
import os
import chromadb
from chromadb.config import Settings
from rag.embeddings import embedding_function
from rag.bm25 import BM25Index
from eval.retrievers import BM25Retriever, DenseRetriever
product = os.environ.get("PRODUCT_NAME", "morpheus")
root = Path(__file__).resolve().parent.parent
col = chromadb.PersistentClient(
path=str(root / "chroma"),
settings=Settings(anonymized_telemetry=False),
).get_collection(f"{product}_docs", embedding_function=embedding_function())
bm = BM25Index(str(root / "bm25" / f"{product}_docs.db"))
dense_r, bm25_r = DenseRetriever(col), BM25Retriever(bm)
except Exception as e:
args.trace_out.parent.mkdir(parents=True, exist_ok=True)
args.trace_out.write_text("")
args.misses_out.write_text("# Eval misses\n\nno index\n")
print(f"no index ({e}); wrote empty trace")
return 0
rows = []
for q in load_queries(args.queries):
expected = [(e["bundle_id"], e["page_id"]) for e in q["expected"]]
dense_pages = dense_r.retrieve(q["query"], k=50)
bm25_pages = bm25_r.retrieve(q["query"], k=50)
ranked = bm25_pages or dense_pages # default retrieval is BM25-first
top1 = ranked[0] if ranked else None
p1 = p_at_1(ranked, expected)
rows.append({
"query": q["query"],
"expected": [list(p) for p in expected],
"hit": bool(p1),
"p_at_1": int(p1),
"dense_rank": first_ranks(dense_pages),
"bm25_rank": first_ranks(bm25_pages),
"top1": list(top1) if top1 else None,
"top1_source": classify_top1_source(top1, set(dense_pages), set(bm25_pages)),
"ranked_pages": [list(p) for p in ranked],
})
args.trace_out.parent.mkdir(parents=True, exist_ok=True)
with open(args.trace_out, "w") as fh:
for row in rows:
fh.write(json.dumps(row) + "\n")
args.misses_out.write_text(render_misses(rows))
print(f"wrote {args.trace_out} ({len(rows)} queries)")
print(f"wrote {args.misses_out}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
+168 -47
View File
@@ -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 dense table chunks.
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)
+23 -2
View File
@@ -22,8 +22,14 @@ import logging
import time
from typing import Any
import httpx
from chromadb import EmbeddingFunction, Documents, Embeddings
try:
import httpx
from chromadb import Documents, EmbeddingFunction, Embeddings
except ImportError:
httpx = None # type: ignore[assignment]
EmbeddingFunction = object # type: ignore[misc,assignment]
Documents = list # type: ignore[misc,assignment]
Embeddings = list # type: ignore[misc,assignment]
log = logging.getLogger(__name__)
@@ -41,6 +47,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
View File
@@ -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)
+71
View File
@@ -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()
+41
View File
@@ -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()