From 169816d98db57a3c5df632f4fd1309328248c004 Mon Sep 17 00:00:00 2001 From: claude Date: Wed, 30 Sep 2026 08:33:16 -0400 Subject: [PATCH] Port docs-mcp-template upgrades (eval, citations, chunking, prefixes) (#13) Co-authored-by: claude --- docs_mcp/format.py | 43 ++++++++ docs_mcp/server.py | 46 +++++---- docs_mcp/test_format.py | 37 +++++++ eval/pvalue.py | 156 +++++++++++++++++++++++++++++ eval/retrievers.py | 4 +- eval/run_eval.py | 46 +++++++-- eval/test_metrics.py | 84 ++++++++++++++++ eval/trace.py | 114 +++++++++++++++++++++ rag/chunk.py | 215 +++++++++++++++++++++++++++++++--------- rag/embeddings.py | 25 ++++- rag/index.py | 7 +- rag/test_chunk.py | 71 +++++++++++++ rag/test_embeddings.py | 41 ++++++++ 13 files changed, 812 insertions(+), 77 deletions(-) create mode 100644 docs_mcp/format.py create mode 100644 docs_mcp/test_format.py create mode 100644 eval/pvalue.py create mode 100644 eval/test_metrics.py create mode 100644 eval/trace.py create mode 100644 rag/test_chunk.py create mode 100644 rag/test_embeddings.py diff --git a/docs_mcp/format.py b/docs_mcp/format.py new file mode 100644 index 0000000..daf01b3 --- /dev/null +++ b/docs_mcp/format.py @@ -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" diff --git a/docs_mcp/server.py b/docs_mcp/server.py index 07a2aa3..94ad3b7 100644 --- a/docs_mcp/server.py +++ b/docs_mcp/server.py @@ -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}" diff --git a/docs_mcp/test_format.py b/docs_mcp/test_format.py new file mode 100644 index 0000000..ee9b75e --- /dev/null +++ b/docs_mcp/test_format.py @@ -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() diff --git a/eval/pvalue.py b/eval/pvalue.py new file mode 100644 index 0000000..cfff50a --- /dev/null +++ b/eval/pvalue.py @@ -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()) diff --git a/eval/retrievers.py b/eval/retrievers.py index 872cf31..230b9b5 100644 --- a/eval/retrievers.py +++ b/eval/retrievers.py @@ -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) diff --git a/eval/run_eval.py b/eval/run_eval.py index 8daa807..c5bc439 100644 --- a/eval/run_eval.py +++ b/eval/run_eval.py @@ -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 diff --git a/eval/test_metrics.py b/eval/test_metrics.py new file mode 100644 index 0000000..882f5b9 --- /dev/null +++ b/eval/test_metrics.py @@ -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() diff --git a/eval/trace.py b/eval/trace.py new file mode 100644 index 0000000..a52e19f --- /dev/null +++ b/eval/trace.py @@ -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()) diff --git a/rag/chunk.py b/rag/chunk.py index c937c1f..f8c4f2b 100644 --- a/rag/chunk.py +++ b/rag/chunk.py @@ -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) diff --git a/rag/embeddings.py b/rag/embeddings.py index c1341e5..0eef5c8 100644 --- a/rag/embeddings.py +++ b/rag/embeddings.py @@ -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): diff --git a/rag/index.py b/rag/index.py index f9b5ce2..6df294b 100644 --- a/rag/index.py +++ b/rag/index.py @@ -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) diff --git a/rag/test_chunk.py b/rag/test_chunk.py new file mode 100644 index 0000000..81213ea --- /dev/null +++ b/rag/test_chunk.py @@ -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() diff --git a/rag/test_embeddings.py b/rag/test_embeddings.py new file mode 100644 index 0000000..abcd1be --- /dev/null +++ b/rag/test_embeddings.py @@ -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()