From d131e93252aacc22b6de59fd862c3bf5e6f966c1 Mon Sep 17 00:00:00 2001 From: claude Date: Wed, 30 Sep 2026 08:22:50 -0400 Subject: [PATCH] feat: citations, eval sidecar/pvalue/trace, nomic prefixes mcp 2.x already on main. Does not replace the variety chunker (one chunk per variety is the anti-hallucination contract). - Numbered [1] citations on search_docs / search_trials - Eval JSONL sidecar, k-curve, eval.pvalue, eval.trace - Nomic prefixes at embed time only; stored text unprefixed Closes #25 --- docs_mcp/format.py | 43 ++++++++++ docs_mcp/server.py | 83 +++++++++---------- docs_mcp/test_format.py | 37 +++++++++ eval/pvalue.py | 156 +++++++++++++++++++++++++++++++++++ eval/results/pre-template.md | 41 +++++++++ eval/retrievers.py | 12 ++- eval/run_eval.py | 42 +++++++++- eval/test_pvalue.py | 25 ++++++ eval/trace.py | 50 +++++++++++ rag/embeddings.py | 25 +++++- rag/index.py | 7 +- rag/test_embeddings.py | 41 +++++++++ 12 files changed, 509 insertions(+), 53 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/results/pre-template.md create mode 100644 eval/test_pvalue.py create mode 100644 eval/trace.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 00000000..daf01b3a --- /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 4b5428b8..ab0b7d75 100644 --- a/docs_mcp/server.py +++ b/docs_mcp/server.py @@ -29,6 +29,7 @@ from mcp.server.mcpserver import MCPServer from mcp.server.transport_security import TransportSecuritySettings from pydantic import Field +from .format import format_search_hit from .usage import TimedCall log = logging.getLogger(__name__) @@ -182,6 +183,17 @@ def _collection(): return _chroma_collection +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_index(): """Return the BM25 index, or None if it doesn't exist on disk.""" global _bm25 @@ -257,10 +269,8 @@ def _read_markdown(source: str, source_key: str) -> str | None: return None -def _format_hit(doc: str, meta: dict, distance: float | None = None) -> str: - """Render one retrieval hit as a fenced markdown block with full - provenance attached. ``doc`` is the chunk text; ``meta`` is the - chunk's metadata dict.""" +def _format_hit(n: int, doc: str, meta: dict, distance: float | None = None) -> str: + """Render one numbered retrieval hit. ``n`` is 1-based display order.""" src_url = meta.get("source_url") or "" src_key = meta.get("source_key") or "" src = meta.get("source") or "" @@ -268,15 +278,10 @@ def _format_hit(doc: str, meta: dict, distance: float | None = None) -> str: brand = meta.get("brand") or "" crop = meta.get("crop") or "" name = meta.get("product_name") or src_key - - header = ( - f"### {name} \n" - f"`{src}::{src_key}` — {vendor} / {brand} / {crop} \n" - f"<{src_url}>" - ) + extra = f"`{src}::{src_key}` — {vendor} / {brand} / {crop}" if distance is not None: - header += f" \n_(distance={distance:.4f})_" - return f"{header}\n\n{doc.strip()}\n" + extra += f" distance={distance:.4f}" + return format_search_hit(n, name, src_url, doc, extra) def _rrf_fuse(rankings: list[list[str]], k: int = RRF_K) -> list[str]: @@ -548,7 +553,8 @@ def search_docs( """Search the seed-variety corpus for hybrids/varieties matching a query. Returns the top-k variety chunks with their full source URLs, ratings, - maturity, traits, and regional listings. Optional filters narrow to one + maturity, traits, and regional listings. Hits are numbered [1]… for + citation. Optional filters narrow to one crop, brand, vendor, or scraper source. Use **list_versions()** first to discover valid facet values. Use **lookup_variety()** on any candidate the user is serious about — that returns the canonical @@ -593,11 +599,7 @@ def search_docs( ) try: - dense = col.query( - query_texts=[query], - n_results=pool_size, - where=where, - ) + dense = _query_dense(col, query, pool_size, where) except Exception as exc: # noqa: BLE001 _call.set(error_dense=str(exc), hits_returned=0) return f"_(retrieval failed: {exc})_" @@ -702,20 +704,21 @@ def search_docs( ) blocks: list[str] = [] - for cid in final_ids: + for n, cid in enumerate(final_ids, start=1): doc = id_to_doc.get(cid, "") meta = id_to_meta.get(cid, {}) dist = id_to_dist.get(cid) if not used_hybrid else None - blocks.append(_format_hit(doc, meta, dist)) + blocks.append(_format_hit(n, doc, meta, dist)) header = ( f"# Search results — {len(final_ids)} variety chunk" f"{'s' if len(final_ids) != 1 else ''}" f"{' (dense + BM25 hybrid)' if used_hybrid else ' (dense only)'}\n" - f"_Use `lookup_variety(source_key=...)` on any candidate " - f"to fact-check ratings from the canonical sidecar._\n\n---\n\n" + f"_Hits are numbered [1]… for citation. " + f"Use `lookup_variety(source_key=...)` on any candidate " + f"to fact-check ratings from the canonical sidecar._\n\n" ) - return header + "\n---\n\n".join(blocks) + return header + "\n\n".join(blocks) + "\n" @mcp.tool() @@ -975,11 +978,7 @@ def search_trials( full_query = f"{query} {product}" try: - dense = col.query( - query_texts=[full_query], - n_results=pool_size, - where=where, - ) + dense = _query_dense(col, full_query, pool_size, where) except Exception as exc: # noqa: BLE001 _call.set(error_dense=str(exc), hits_returned=0) return f"_(trial retrieval failed: {exc})_" @@ -1093,28 +1092,27 @@ def search_trials( ) blocks: list[str] = [] - for cid in final_ids: + for n, cid in enumerate(final_ids, start=1): doc = id_to_doc.get(cid, "") meta = id_to_meta.get(cid, {}) dist = id_to_dist.get(cid) if not used_hybrid else None - blocks.append(_format_trial_hit(doc, meta, dist)) + blocks.append(_format_trial_hit(n, doc, meta, dist)) header = ( f"# Trial search results — {len(final_ids)} trial document" f"{'s' if len(final_ids) != 1 else ''}" f"{' (dense + BM25 hybrid)' if used_hybrid else ' (dense only)'}\n" - f"_Use `get_page(source=..., source_key=...)` to read the " + f"_Hits are numbered [1]… for citation. " + f"Use `get_page(source=..., source_key=...)` to read the " f"full trial body. Use `lookup_variety(source_key=...)` on " f"any product code to verify its identity (RM, traits, " - f"disease ratings)._\n\n---\n\n" + f"disease ratings)._\n\n" ) - return header + "\n---\n\n".join(blocks) + return header + "\n\n".join(blocks) + "\n" -def _format_trial_hit(doc: str, meta: dict, distance: float | None = None) -> str: - """Trial-specific result header. Highlights crop/state/year and - sources URL (vs variety hits which emphasize brand + product - identity).""" +def _format_trial_hit(n: int, doc: str, meta: dict, distance: float | None = None) -> str: + """Trial-specific numbered hit. Highlights crop/state/year.""" src_url = meta.get("source_url") or "" src_key = meta.get("source_key") or "" src = meta.get("source") or "" @@ -1125,15 +1123,10 @@ def _format_trial_hit(doc: str, meta: dict, distance: float | None = None) -> st title_bits = [b for b in [crop.title(), region or state, str(year) if year else ""] if b] title = " · ".join(title_bits) if title_bits else src_key - - header = ( - f"### Trial: {title} \n" - f"`{src}::{src_key}` — {meta.get('vendor', '')} / {meta.get('brand', '')} \n" - f"<{src_url}>" - ) + extra = f"`{src}::{src_key}` — {meta.get('vendor', '')} / {meta.get('brand', '')}" if distance is not None: - header += f" \n_(distance={distance:.4f})_" - return f"{header}\n\n{doc.strip()}\n" + extra += f" distance={distance:.4f}" + return format_search_hit(n, f"Trial: {title}", src_url, doc, extra) @mcp.tool() diff --git a/docs_mcp/test_format.py b/docs_mcp/test_format.py new file mode 100644 index 00000000..ee9b75ed --- /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 00000000..cfff50ac --- /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/results/pre-template.md b/eval/results/pre-template.md new file mode 100644 index 00000000..1d23700c --- /dev/null +++ b/eval/results/pre-template.md @@ -0,0 +1,41 @@ +# seed-mcp retrieval eval — k=5 + +_21 golden queries × 4 retrievers_ + +## Summary + +| Retriever | Passed | Recall | P@1 | MRR | Avg ms | +|---|---|---|---|---|---| +| **hybrid+rerank** | 21/21 | 100.00% | 90.48% | 0.905 | 2064 | +| **bm25** | 20/21 | 95.24% | 80.95% | 0.833 | 5 | +| **hybrid** | 15/21 | 71.43% | 61.90% | 0.619 | 73 | +| **dense** | 14/21 | 66.67% | 38.10% | 0.440 | 79 | + +**Recall** = % of queries where ≥1 top-k chunk satisfied the spec. **P@1** = % where the very first result satisfied it. **MRR** = mean of `1 / rank-of-first-satisfying-result` (0 if missed). + +## Per-query results + +| Query | bm25 | dense | hybrid | hybrid+rerank | +|---|---|---|---|---| +| `DKC62-08RIB ratings` | ✅ #1 | ❌ | ❌ | ✅ #1 | +| `AG29XF4 disease ratings` | ✅ #1 | ❌ | ❌ | ✅ #1 | +| `WB6430 westbred wheat` | ✅ #1 | ❌ | ❌ | ✅ #1 | +| `E085Z5 corn` | ✅ #1 | ❌ | ❌ | ✅ #1 | +| `AP Iliad wheat performance` | ✅ #1 | ❌ | ❌ | ✅ #1 | +| `drought tolerant corn for sandy soil short season Iowa` | ✅ #2 | ✅ #1 | ✅ #1 | ✅ #1 | +| `soybean cyst nematode SCN resistant variety` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 | +| `Phytophthora resistance Rps3a soybean` | ✅ #1 | ✅ #2 | ✅ #1 | ✅ #1 | +| `XtendFlex soybean Northern Plains` | ❌ | ✅ #1 | ✅ #1 | ✅ #1 | +| `Hard Red Spring wheat stripe rust resistance` | ✅ #1 | ✅ #3 | ✅ #1 | ✅ #1 | +| `Soft White Winter wheat Pacific Northwest` | ✅ #1 | ✅ #5 | ✅ #1 | ✅ #1 | +| `Goss's Wilt resistance corn` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 | +| `best corn 2024 Iowa` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 | +| `Indiana corn yield comparison 2024` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 | +| `AP Iliad Idaho wheat trial` | ✅ #1 | ✅ #5 | ✅ #1 | ✅ #1 | +| `DKC65-95 corn yield in trials` | ✅ #1 | ❌ | ✅ #1 | ✅ #1 | +| `NK1701 corn trials head to head` | ✅ #1 | ❌ | ❌ | ✅ #1 | +| `silage corn high milk per acre dairy` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 | +| `soybean 2025 Minnesota top performers` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 | +| `Pioneer P1142 hybrid recommendation` | ✅ | ✅ | ✅ | ✅ | +| `DKC65-20 yield Alabama trial` | ✅ | ✅ | ✅ | ✅ | + diff --git a/eval/retrievers.py b/eval/retrievers.py index d176254e..e2d38d1e 100644 --- a/eval/retrievers.py +++ b/eval/retrievers.py @@ -71,8 +71,10 @@ class DenseRetriever: def retrieve(self, query: str, k: int, filters: dict | None) -> list[str]: where = _build_where(filters) try: + from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts + qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0] r = self.col.query( - query_texts=[query], n_results=max(k, self.pool), where=where, + query_embeddings=[qvec], n_results=max(k, self.pool), where=where, ) except Exception: return [] @@ -105,7 +107,9 @@ class HybridRetriever: def retrieve(self, query: str, k: int, filters: dict | None) -> list[str]: where = _build_where(filters) try: - d = self.col.query(query_texts=[query], n_results=self.pool, where=where) + from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts + qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0] + d = self.col.query(query_embeddings=[qvec], n_results=self.pool, where=where) dense_ids = (d.get("ids") or [[]])[0] except Exception: dense_ids = [] @@ -138,8 +142,10 @@ class HybridRerankRetriever: def retrieve(self, query: str, k: int, filters: dict | None) -> list[str]: where = _build_where(filters) try: + from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts + qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0] d = self.col.query( - query_texts=[query], n_results=self.pool, where=where, + query_embeddings=[qvec], n_results=self.pool, where=where, include=["documents"], ) dense_ids = (d.get("ids") or [[]])[0] diff --git a/eval/run_eval.py b/eval/run_eval.py index 9c8f40d1..f4a8443b 100644 --- a/eval/run_eval.py +++ b/eval/run_eval.py @@ -265,10 +265,15 @@ 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) + p.add_argument("--trace", action="store_true") p.add_argument("--rerank-url", default=os.environ.get("RERANK_URL", "")) p.add_argument("--product-name", default=os.environ.get("PRODUCT_NAME", "crop_seed")) 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}") @@ -300,14 +305,40 @@ def main() -> int: for r in retrievers: print(f"running {r.name}...") for q in queries: - res = _evaluate_one(r, q, args.k, col) + res = _evaluate_one(r, q, max_k, col) all_results.append(res) summary = _aggregate(all_results) md = _emit_markdown(queries, all_results, summary, args.k) + # k-curve from rank_first_match at max_k + md += "\n## k-curve (P@1 / recall from rank_first_match)\n\n" + md += "| Retriever | " + " | ".join(f"P@1@k={k}" for k in ks) + " |\n" + md += "|" + "---|" * (len(ks) + 1) + "\n" + by_r: dict[str, list[dict]] = {} + for row in all_results: + by_r.setdefault(row["retriever"], []).append(row) + for name, rows in by_r.items(): + cells = [] + for kk in ks: + hits = sum(1 for r in rows if r.get("rank_first_match") and r["rank_first_match"] <= kk) + cells.append(f"{hits / len(rows):.3f}" if rows else "0") + md += f"| `{name}` | " + " | ".join(cells) + " |\n" args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(md, encoding="utf-8") + sidecar = args.output.with_suffix(".jsonl") + with open(sidecar, "w") as fh: + for r in all_results: + rank = r.get("rank_first_match") + fh.write(json.dumps({ + "query": r["query"], + "retriever": r["retriever"], + "passed": bool(r.get("passed")), + "p_at_1": 1 if rank == 1 else 0, + "rr": (1.0 / rank) if rank else 0.0, + "rank_first_match": rank, + }) + "\n") print(f"\nreport: {args.output}") + print(f"sidecar: {sidecar}") print() # Print summary to stdout too for line in md.split("\n"): @@ -315,6 +346,15 @@ def main() -> int: print(line) if line.startswith("## Per-query"): break + if args.compare: + from eval.pvalue import compare, load_sidecar, render + print() + print(render(compare(load_sidecar(sidecar), load_sidecar(args.compare))), end="") + if args.trace: + from eval.trace import render_misses + misses_path = args.output.with_name("misses.md") + misses_path.write_text(render_misses(all_results)) + print(f"misses: {misses_path}") return 0 diff --git a/eval/test_pvalue.py b/eval/test_pvalue.py new file mode 100644 index 00000000..b8288d7c --- /dev/null +++ b/eval/test_pvalue.py @@ -0,0 +1,25 @@ +"""Stdlib tests for the permutation test. No Chroma.""" +from __future__ import annotations + +import unittest + +from eval.pvalue import paired_permutation + + +class PermutationTests(unittest.TestCase): + def test_identical_is_not_significant(self) -> None: + a = [1.0, 0.5, 0.0, 1.0] + r = paired_permutation(a, list(a), n_resamples=200, seed=0) + self.assertEqual(r["Diff(A-B)"], 0.0) + self.assertFalse(r["significant"]) + + def test_large_shift_is_significant(self) -> None: + a = [1.0] * 20 + b = [0.0] * 20 + r = paired_permutation(a, b, n_resamples=500, seed=0) + self.assertTrue(r["significant"]) + self.assertLess(r["p_value"], 0.05) + + +if __name__ == "__main__": + unittest.main() diff --git a/eval/trace.py b/eval/trace.py new file mode 100644 index 00000000..cbea7438 --- /dev/null +++ b/eval/trace.py @@ -0,0 +1,50 @@ +"""Miss dump from an eval sidecar or in-memory rows. + + python -m eval.trace --sidecar eval/results/baseline.jsonl +""" +from __future__ import annotations + +import argparse +import json +from pathlib import Path + + +def render_misses(rows: list[dict]) -> str: + misses = [r for r in rows if not r.get("passed") and r.get("p_at_1", 1) == 0] + # passed=False is the seed schema; p_at_1==0 covers pvalue sidecars + if not misses: + misses = [r for r in rows if not r.get("passed")] + if not misses: + return "# Eval misses\n\n_(none)_\n" + lines = [f"# Eval misses ({len(misses)})", ""] + for row in misses: + lines += [ + f"## {row.get('query', '')}", + "", + f"- retriever: `{row.get('retriever')}`", + f"- rank_first_match: `{row.get('rank_first_match')}`", + f"- kind: `{row.get('kind', '')}`", + "", + ] + return "\n".join(lines) + + +def main() -> int: + p = argparse.ArgumentParser() + p.add_argument("--sidecar", type=Path, default=Path("eval/results/baseline.jsonl")) + p.add_argument("--misses-out", type=Path, default=Path("eval/results/misses.md")) + args = p.parse_args() + if not args.sidecar.exists(): + args.misses_out.parent.mkdir(parents=True, exist_ok=True) + args.misses_out.write_text("# Eval misses\n\nno sidecar\n") + print(f"no sidecar ({args.sidecar}); wrote empty misses") + return 0 + rows = [json.loads(line) for line in args.sidecar.read_text().splitlines() if line.strip()] + args.misses_out.parent.mkdir(parents=True, exist_ok=True) + args.misses_out.write_text(render_misses(rows)) + print(f"wrote {args.misses_out} ({len(rows)} rows)") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/rag/embeddings.py b/rag/embeddings.py index 84d3bbdb..65810f7a 100644 --- a/rag/embeddings.py +++ b/rag/embeddings.py @@ -14,8 +14,14 @@ import os import logging from typing import Any -import httpx -from chromadb import EmbeddingFunction, Documents, Embeddings +try: + import httpx + from chromadb import Documents, EmbeddingFunction, Embeddings +except ImportError: # unit tests for apply_prefix / embed_texts + 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__) @@ -23,6 +29,21 @@ OLLAMA_URLS = [u.strip() for u in os.environ.get("OLLAMA_URL", "http://localhost:11434").split(",") if u.strip()] 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 a453b5ba..a2544992 100644 --- a/rag/index.py +++ b/rag/index.py @@ -23,7 +23,7 @@ import chromadb from chromadb.config import Settings from .chunk import chunks_from_variety, chunks_from_trial -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") @@ -83,9 +83,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_embeddings.py b/rag/test_embeddings.py new file mode 100644 index 00000000..abcd1be4 --- /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()