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

Merged
claude merged 1 commits from issue-15 into main 2026-09-29 22:27:17 -04:00
13 changed files with 803 additions and 74 deletions
Showing only changes of commit 205fcfd7ef - 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 mcp.server.transport_security import TransportSecuritySettings
from pydantic import Field from pydantic import Field
from .format import format_search_hits
from .usage import TimedCall from .usage import TimedCall
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@@ -145,6 +146,17 @@ def _collection():
return _CHROMA 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(): def _bm25():
"""Lazy BM25Index handle. None if the FTS5 db isn't built.""" """Lazy BM25Index handle. None if the FTS5 db isn't built."""
global _BM25 global _BM25
@@ -258,9 +270,10 @@ def search_docs(
"""Search the HPE Morpheus VM Essentials (HVM) docs corpus. """Search the HPE Morpheus VM Essentials (HVM) docs corpus.
Returns the top-k most relevant chunks (with full source page URLs) Returns the top-k most relevant chunks (with full source page URLs)
given a natural-language query. Optional filters narrow the search given a natural-language query. Hits are numbered [1]… for citation.
to one version, one platform, or one bundle. Use list_versions() Optional filters narrow the search to one version, one platform, or
first if you need to discover the available facet values. 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 Call this tool whenever the user asks anything that should be
answerable from the official product documentation — install, answerable from the official product documentation — install,
@@ -296,7 +309,7 @@ def search_docs(
if HYBRID_SEARCH and bm is not None: if HYBRID_SEARCH and bm is not None:
try: 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] dense_ids = (dense_res.get("ids") or [[]])[0]
bm_hits = bm.query(query, n=pool, where=bm25_where) bm_hits = bm.query(query, n=pool, where=bm25_where)
bm_ids = [cid for cid, _s in bm_hits] 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) log.warning("BM25 retrieval failed, falling back to dense: %s", e)
if not docs: 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] docs = (res.get("documents") or [[]])[0]
metas = (res.get("metadatas") or [[]])[0] metas = (res.get("metadatas") or [[]])[0]
dists = (res.get("distances") 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 = bm.query(query, n=pool_size, where=bm25_where) if bm else []
extra_ids = [cid for cid, _s in extra] extra_ids = [cid for cid, _s in extra]
else: 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] extra_ids = (extra_res.get("ids") or [[]])[0]
if extra_ids: if extra_ids:
d2, m2, _ = _enrich_from_chroma(col, extra_ids, None) d2, m2, _ = _enrich_from_chroma(col, extra_ids, None)
@@ -371,22 +384,21 @@ def search_docs(
if not docs: if not docs:
return f"_No matches for `{query}`._" 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): for doc, meta, dist in zip(docs, metas, dists):
bid = meta.get("bundle_id", "") bid = meta.get("bundle_id", "")
pid = meta.get("page_id", "") pid = meta.get("page_id", "")
title = meta.get("title") or pid title = meta.get("title") or pid
ver = meta.get("version") or "" ver = meta.get("version") or ""
url = _source_url(bid, pid) url = _source_url(bid, pid)
header = f"## {title}" extra = f"score={1 - dist:.3f}"
if ver: if ver:
header += f" _(v{ver})_" extra = f"v{ver} " + extra
out.append(header) if pid:
out.append(f"[{bid}/{pid}]({url}) · score={1 - dist:.3f}") extra += f" page_id: `{pid}`"
out.append("") hits.append((title, url, doc, extra))
out.append(doc.strip()) header = f"# {len(docs)} result(s) for `{query}`\n\n"
out.append("") return header + format_search_hits(hits)
return "\n".join(out)
@mcp.tool() @mcp.tool()
@@ -1073,8 +1085,8 @@ def find_doc_inconsistencies(
return f"Couldn't open Chroma collection: {e}" return f"Couldn't open Chroma collection: {e}"
where = _build_where(version, platform, bundle_id) where = _build_where(version, platform, bundle_id)
try: try:
res = col.query(query_texts=[scope_query], n_results=max_pages * 3, res = _query_dense(col, scope_query, max_pages * 3, where,
where=where, include=["metadatas"]) include=["metadatas"])
except Exception as e: except Exception as e:
_call.set(error=f"query: {e}") _call.set(error=f"query: {e}")
return f"Scope query failed: {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 self.pool = pool
def retrieve(self, query: str, k: int = 10) -> list[tuple[str, str]]: 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] ids = (res.get("ids") or [[]])[0]
return _collapse_to_pages(ids, k) return _collapse_to_pages(ids, k)
+37 -7
View File
@@ -34,6 +34,12 @@ def load_queries(path: Path) -> list[dict]:
return [json.loads(line) for line in fh if line.strip()] 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: def reciprocal_rank(retrieved: list[tuple[str, str]], expected: list[tuple[str, str]]) -> float:
expected_set = set(expected) expected_set = set(expected)
for i, page in enumerate(retrieved, start=1): for i, page in enumerate(retrieved, start=1):
@@ -65,8 +71,12 @@ def main() -> int:
p = argparse.ArgumentParser() p = argparse.ArgumentParser()
p.add_argument("--queries", type=Path, default=Path("eval/queries.jsonl")) p.add_argument("--queries", type=Path, default=Path("eval/queries.jsonl"))
p.add_argument("--k", type=int, default=5) 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("--output", type=Path, default=Path("eval/results/baseline.md"))
p.add_argument("--compare", type=Path, default=None)
args = p.parse_args() 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(): if not args.queries.exists():
print(f"queries file not found: {args.queries}") print(f"queries file not found: {args.queries}")
@@ -109,34 +119,39 @@ def main() -> int:
rows: dict[str, dict[str, float]] = {} rows: dict[str, dict[str, float]] = {}
per_query: list[dict] = [] per_query: list[dict] = []
for r in retrievers: 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 elapsed_sum = 0.0
for q in queries: for q in queries:
expected = [(e["bundle_id"], e["page_id"]) for e in q["expected"]] expected = [(e["bundle_id"], e["page_id"]) for e in q["expected"]]
t0 = time.time() 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 elapsed = time.time() - t0
mrr = reciprocal_rank(retrieved, expected) mrr = reciprocal_rank(retrieved, expected)
p1 = p_at_1(retrieved, expected)
recall = recall_at_k(retrieved, expected, args.k) recall = recall_at_k(retrieved, expected, args.k)
ndcg = ndcg_at_k(retrieved, expected, args.k) ndcg = ndcg_at_k(retrieved, expected, args.k)
mrr_sum += mrr mrr_sum += mrr
p1_sum += p1
recall_sum += recall recall_sum += recall
ndcg_sum += ndcg ndcg_sum += ndcg
elapsed_sum += elapsed elapsed_sum += elapsed
per_query.append({ per_query.append({
"retriever": r.name, "query": q["query"], "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, "top1": list(retrieved[0]) if retrieved else None,
"ranked": [list(p) for p in retrieved],
"elapsed_s": round(elapsed, 3), "elapsed_s": round(elapsed, 3),
}) })
n = len(queries) n = len(queries)
rows[r.name] = { rows[r.name] = {
"P@1": p1_sum / n,
"MRR": mrr_sum / n, "MRR": mrr_sum / n,
f"Recall@{args.k}": recall_sum / n, f"Recall@{args.k}": recall_sum / n,
f"nDCG@{args.k}": ndcg_sum / n, f"nDCG@{args.k}": ndcg_sum / n,
"avg_latency_s": elapsed_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"Recall@{args.k}={rows[r.name][f'Recall@{args.k}']:.3f} "
f"nDCG@{args.k}={rows[r.name][f'nDCG@{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") 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) args.output.parent.mkdir(parents=True, exist_ok=True)
md = [f"# Retrieval eval — k={args.k}", "", md = [f"# Retrieval eval — k={args.k}", "",
f"_{len(queries)} hand-curated queries, generated {time.strftime('%Y-%m-%d %H:%M:%S')}_", "", 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(): 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 |") f"| {m[f'nDCG@{args.k}']:.3f} | {m['avg_latency_s']*1000:.0f}ms |")
md += ["", "## Per-query results", "", md += ["", "## Per-query results", "",
"| Retriever | Query | MRR | top-1 |", "| --- | --- | ---: | --- |"] "| 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 "—" 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} |") md.append(f"| `{r['retriever']}` | {r['query'][:60]} | {r['mrr']:.3f} | {top1} |")
args.output.write_text("\n".join(md) + "\n") 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 {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 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", "hvm")
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 # HVM 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 Chunk by semantic section (ATX headings), not raw page/length. A
significantly from prose. The output shape (id, text, metadata) is synthetic chunk 0 (title + first paragraph + optional keyword bag) is
fixed by the downstream Chroma + BM25 indexing in rag/index.py — don't always emitted first — dense retrieval lands on it. Do not drop it.
change that.
The key knob you'll tune per product is chunk-0. Dense retrieval lands No chonkie / langchain dependency. The output shape (id, text, metadata)
on chunk 0 first for most queries. Make it a synthetic chunk built is fixed by rag/index.py — don't change that. `heading_path` is optional
from: metadata (e.g. "Install > Linux") and is not required for indexing.
- 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.
""" """
from __future__ import annotations from __future__ import annotations
@@ -26,17 +14,12 @@ import re
from typing import Iterator 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 CHARS_PER_TOKEN = 4
TARGET_TOKENS = 500 TARGET_TOKENS = 500
TARGET_CHARS = TARGET_TOKENS * CHARS_PER_TOKEN TARGET_CHARS = TARGET_TOKENS * CHARS_PER_TOKEN
# Hard cap: nomic-embed-text's context is 2048 tokens. Anything larger # nomic-embed-text context is 2048 tokens. Markdown tables with lots of
# 400s the entire embed batch. 6000 chars works for prose but markdown # `|` tokenize ~1.4× denser than prose; 4000 chars stays under 2048 even
# tables with lots of `|` separators tokenize ~1.4× denser; a 5839-char # for qualification-matrix chunks (a 5839-char table crashed a rebuild).
# 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.
MAX_CHARS = 4000 MAX_CHARS = 4000
@@ -97,6 +80,158 @@ def split_paragraphs(md: str) -> list[str]:
return [b for b in blocks if b] 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( def chunks_from_page(
text: str, text: str,
page_id: str, page_id: str,
@@ -112,7 +247,6 @@ def chunks_from_page(
if not paragraphs: if not paragraphs:
return return
# ----- Chunk 0: synthetic anchor for dense retrieval ---------
title = metadata.get("title") or page_id title = metadata.get("title") or page_id
first_para = next((p for p in paragraphs if not p.startswith("#")), "") first_para = next((p for p in paragraphs if not p.startswith("#")), "")
chunk0_body = ( chunk0_body = (
@@ -127,28 +261,15 @@ def chunks_from_page(
"metadata": {**metadata, "ordinal": 0}, "metadata": {**metadata, "ordinal": 0},
} }
# ----- Body chunks: pack paragraphs up to TARGET_CHARS -------
ordinal = 1 ordinal = 1
for body, heading_path in _pack_recursive(paragraphs, []):
def emit(buf: list[str]) -> Iterator[dict]: for piece in _hard_split(body):
nonlocal ordinal meta = {**metadata, "ordinal": ordinal}
merged = "\n\n".join(buf) if heading_path:
for piece in _hard_split(merged): meta["heading_path"] = heading_path
yield { yield {
"id": f"{metadata['bundle_id']}::{page_id}::{ordinal}", "id": f"{metadata['bundle_id']}::{page_id}::{ordinal}",
"text": piece, "text": piece,
"metadata": {**metadata, "ordinal": ordinal}, "metadata": meta,
} }
ordinal += 1 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)
+15
View File
@@ -41,6 +41,21 @@ def _resolve_urls() -> list[str]:
OLLAMA_URLS = _resolve_urls() OLLAMA_URLS = _resolve_urls()
EMBED_MODEL = os.environ.get("EMBED_MODEL", "nomic-embed-text") EMBED_MODEL = os.environ.get("EMBED_MODEL", "nomic-embed-text")
EMBED_DIM = int(os.environ.get("EMBED_DIM", "768")) 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): class OllamaEmbeddings(EmbeddingFunction):
+5 -2
View File
@@ -18,7 +18,7 @@ import chromadb
from chromadb.config import Settings from chromadb.config import Settings
from .chunk import chunks_from_page 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__) log = logging.getLogger(__name__)
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s") logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
@@ -78,9 +78,12 @@ def upsert_to_chroma(records: list[dict]) -> int:
total = 0 total = 0
for i in range(0, len(records), BATCH): for i in range(0, len(records), BATCH):
chunk = records[i:i + 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( col.upsert(
ids=[r["id"] for r in chunk], 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], metadatas=[r["metadata"] for r in chunk],
) )
total += len(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()