5 Commits
17 changed files with 878 additions and 108 deletions
+11 -11
View File
@@ -17,7 +17,7 @@ once deployed.
| Tool | Use | | Tool | Use |
|---|---| |---|---|
| `search_docs` | BM25-default search with optional version / platform / bundle filters; cross-encoder reranked when `RERANK_URL` is set | | `search_docs` | Dense-default search with optional version / platform / bundle filters; rerank only when `RERANK_ENABLED=true` |
| `get_page` | Full markdown of one page with metadata header + source URL | | `get_page` | Full markdown of one page with metadata header + source URL |
| `list_versions` | Discover available versions, doc types, and bundle slugs | | `list_versions` | Discover available versions, doc types, and bundle slugs |
| `list_cluster` | Cross-version peers of a page (synthesized from same-GUID overlap) | | `list_cluster` | Cross-version peers of a page (synthesized from same-GUID overlap) |
@@ -51,19 +51,19 @@ peer mapping is free (no fuzzy matching needed).
## Retrieval ## Retrieval
Eval against 22 hand-curated golden queries — see Eval against 22 hand-curated golden queries. May 2026 baseline is
[`eval/results/baseline.md`](eval/results/baseline.md): [`eval/results/baseline.md`](eval/results/baseline.md); after nomic
prefixes + heading-recursive chunking (2026-09-30, live index):
| Retriever | MRR | Recall@5 | nDCG@5 | latency | | Retriever | P@1 | MRR | Recall@5 | nDCG@5 |
|---|---:|---:|---:|---:| |---|---:|---:|---:|---:|
| dense (Ollama nomic-embed-text) | 0.539 | 0.621 | 0.558 | 88 ms | | **dense** (nomic prefixes) | **0.955** | **0.966** | **0.985** | **0.972** |
| BM25 (SQLite FTS5) | 0.880 | 0.909 | 0.883 | 3 ms | | hybrid (dense + BM25 + RRF) | 0.955 | 0.961 | 0.955 | 0.955 |
| hybrid (dense + BM25 + RRF) | 0.692 | 0.818 | 0.713 | 69 ms | | BM25 (SQLite FTS5) | 0.864 | 0.882 | 0.909 | 0.886 |
| **bm25 + jina-rerank** | **0.920** | **0.939** | **0.927** | 490 ms (CPU) / ~50 ms (GPU) | | bm25 + jina-rerank | 0.773 | 0.827 | 0.848 | 0.816 |
HPE docs use controlled vocabulary, so lexical match dominates; the `search_docs` defaults to dense. Rerank is opt-in (`RERANK_ENABLED=true`)
cross-encoder cleans up the long tail. See PLAN.md Phase 7/8 for the — it now hurts. Full table: [`eval/results/post-prefix.md`](eval/results/post-prefix.md).
reasoning.
## Architecture ## Architecture
File diff suppressed because one or more lines are too long
+6 -6
View File
@@ -32,16 +32,16 @@ services:
# add that hostname here. "*" disables the rebind check entirely. # add that hostname here. "*" disables the rebind check entirely.
MCP_ALLOWED_HOSTS: "hvm-docs-mcp,localhost,127.0.0.1" MCP_ALLOWED_HOSTS: "hvm-docs-mcp,localhost,127.0.0.1"
# Phase 6 — reranker sidecar (jina-reranker-v2-base via llama.cpp). # Phase 6 — reranker sidecar is wired but off. Eval 2026-09-30
# (post nomic prefixes): dense MRR=0.966 vs bm25+rerank 0.827.
# Set RERANK_ENABLED=true to turn the sidecar back on.
RERANK_URL: http://hvm-rerank:8080 RERANK_URL: http://hvm-rerank:8080
RERANK_POOL: "200" RERANK_POOL: "200"
RERANK_TIMEOUT: "30" RERANK_TIMEOUT: "30"
RERANK_ENABLED: "false"
# Phase 8 — hybrid retrieval (BM25 + dense + RRF). # Phase 8 — hybrid retrieval (BM25 + dense + RRF). Off: dense is
# Eval on the HVM corpus (eval/results/baseline.md, 2026-05-22) shows # the default (eval/results/post-prefix.md). BM25 is fallback only.
# BM25-default + reranker beats hybrid on every metric (MRR 0.920 vs
# 0.875). Leaving HYBRID_SEARCH off so search_docs runs BM25-first +
# reranker; dense is the fallback when BM25 finds nothing.
HYBRID_SEARCH: "false" HYBRID_SEARCH: "false"
# Phase 10 — usage telemetry. # Phase 10 — usage telemetry.
+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"
+66 -32
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__)
@@ -60,6 +61,9 @@ API_LESSONS_MD = Path(__file__).resolve().parent / "api_lessons.md"
RERANK_URL = os.environ.get("RERANK_URL", "").rstrip("/") or None RERANK_URL = os.environ.get("RERANK_URL", "").rstrip("/") or None
RERANK_POOL = int(os.environ.get("RERANK_POOL", "50")) RERANK_POOL = int(os.environ.get("RERANK_POOL", "50"))
RERANK_TIMEOUT = float(os.environ.get("RERANK_TIMEOUT", "30")) RERANK_TIMEOUT = float(os.environ.get("RERANK_TIMEOUT", "30"))
# Opt-in. Watchtower keeps the old container env (RERANK_URL is set in
# live compose); default off so a code-only ship actually stops reranking.
RERANK_ENABLED = os.environ.get("RERANK_ENABLED", "").lower() in ("true", "1", "yes", "on")
HYBRID_SEARCH = os.environ.get("HYBRID_SEARCH", "").lower() in ("true", "1", "yes", "on") HYBRID_SEARCH = os.environ.get("HYBRID_SEARCH", "").lower() in ("true", "1", "yes", "on")
RRF_K = int(os.environ.get("RRF_K", "60")) RRF_K = int(os.environ.get("RRF_K", "60"))
@@ -145,6 +149,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 +273,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,
@@ -282,11 +298,10 @@ def search_docs(
bm25_where = _where_for_bm25(version, platform, bundle_id) bm25_where = _where_for_bm25(version, platform, bundle_id)
pool = max(k * 5, 50) pool = max(k * 5, 50)
# Retrieval mode selection. Eval on this corpus (2026-05-22, 22 golden # Retrieval mode. Eval 2026-09-30 (22 queries, post nomic prefixes +
# queries) showed BM25 MRR=0.88 vs dense MRR=0.54 vs hybrid MRR=0.69 — # heading-recursive chunking): dense MRR=0.966 P@1=0.955 vs
# HPE structured docs use controlled vocabulary, so lexical match wins. # bm25+rerank MRR=0.827 P@1=0.773. Default is dense. HYBRID_SEARCH
# Dense is kept as fallback when BM25 has no tokens to chew on (e.g. # still forces RRF. Rerank is opt-in (RERANK_ENABLED) — it now hurts.
# purely stopword queries). HYBRID_SEARCH=true forces RRF fusion.
bm = _bm25() bm = _bm25()
docs: list[str] = [] docs: list[str] = []
metas: list[dict] = [] metas: list[dict] = []
@@ -296,7 +311,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]
@@ -309,7 +324,18 @@ def search_docs(
else "dense_only") else "dense_only")
retrieval_mode = "hybrid" retrieval_mode = "hybrid"
except Exception as e: except Exception as e:
log.warning("hybrid failed, falling back to BM25→dense: %s", e) log.warning("hybrid failed, falling back to dense: %s", e)
if not docs:
try:
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]
retrieval_mode = "dense"
top1_source = "dense_only"
except Exception as e:
log.warning("dense retrieval failed, falling back to BM25: %s", e)
if not docs and bm is not None: if not docs and bm is not None:
try: try:
@@ -323,16 +349,10 @@ def search_docs(
retrieval_mode = "bm25" retrieval_mode = "bm25"
top1_source = "bm25_only" top1_source = "bm25_only"
except Exception as e: except Exception as e:
log.warning("BM25 retrieval failed, falling back to dense: %s", e) log.warning("BM25 retrieval failed: %s", e)
if not docs:
res = col.query(query_texts=[query], n_results=k, where=where)
docs = (res.get("documents") or [[]])[0]
metas = (res.get("metadatas") or [[]])[0]
dists = (res.get("distances") or [[]])[0]
reranker_fired = False reranker_fired = False
if RERANK_URL and docs: if RERANK_URL and RERANK_ENABLED and docs:
# Pull a deeper pool to give the reranker something to chew on. # Pull a deeper pool to give the reranker something to chew on.
# We over-fetch up to RERANK_POOL chunks from whichever retriever # We over-fetch up to RERANK_POOL chunks from whichever retriever
# already won, then ask the reranker to pick the final top-k. # already won, then ask the reranker to pick the final top-k.
@@ -342,7 +362,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 +391,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()
@@ -810,6 +829,21 @@ def weekly_digest(
return "\n".join(lines) return "\n".join(lines)
def _pub_sort_key(pub: str) -> tuple[int, int]:
"""Sortable (year, month) for an HPE "Published" value like "July 2026".
Plain string comparison ranks "March 2026" above "July 2026".
Unparseable values sort last.
"""
for fmt in ("%B %Y", "%b %Y", "%B %d, %Y", "%Y-%m-%d"):
try:
d = _dt.datetime.strptime(pub.strip(), fmt)
return (d.year, d.month)
except ValueError:
continue
return (0, 0)
@mcp.tool() @mcp.tool()
def corpus_status() -> str: def corpus_status() -> str:
"""Freshness + size of the knowledge base. """Freshness + size of the knowledge base.
@@ -831,7 +865,7 @@ def corpus_status() -> str:
for slug, b in cat.items(): for slug, b in cat.items():
pub = (b.get("dates") or {}).get("Published") pub = (b.get("dates") or {}).get("Published")
if pub: if pub:
if latest_pub is None or pub > latest_pub: if latest_pub is None or _pub_sort_key(pub) > _pub_sort_key(latest_pub):
latest_pub = pub latest_pub = pub
per_bundle.append((slug, pub)) per_bundle.append((slug, pub))
if latest_pub: if latest_pub:
@@ -850,7 +884,7 @@ def corpus_status() -> str:
"", "",
] ]
if per_bundle: if per_bundle:
per_bundle.sort(key=lambda kv: kv[1], reverse=True) per_bundle.sort(key=lambda kv: _pub_sort_key(kv[1]), reverse=True)
lines.append("## Most-recently-edited bundles (by HPE)") lines.append("## Most-recently-edited bundles (by HPE)")
for slug, when in per_bundle[:5]: for slug, when in per_bundle[:5]:
b = cat.get(slug, {}) b = cat.get(slug, {})
@@ -1058,8 +1092,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())
+21
View File
@@ -0,0 +1,21 @@
# Retrieval eval — k=5
_22 hand-curated queries, generated 2026-09-30 02:40:07_
Live container after nomic prefixes + heading-recursive chunking
(`hvm-docs` image `b200c5f76c56`, Watchtower 2026-09-30T02:36Z).
Compared with [`baseline.md`](baseline.md) (2026-05-22, pre-prefix).
| Retriever | P@1 | MRR | Recall@5 | nDCG@5 | avg latency |
| --- | ---: | ---: | ---: | ---: | ---: |
| `dense` | 0.955 | 0.966 | 0.985 | 0.972 | 183ms |
| `bm25` | 0.864 | 0.882 | 0.909 | 0.886 | 8ms |
| `hybrid_rrf` | 0.955 | 0.961 | 0.955 | 0.955 | 113ms |
| `bm25+rerank` | 0.773 | 0.827 | 0.848 | 0.816 | 162ms |
| `hybrid_rrf+rerank` | 0.727 | 0.813 | 0.894 | 0.821 | 254ms |
Dense flipped from worst (May MRR 0.539) to best. Rerank now hurts.
`search_docs` therefore defaults to dense; rerank is opt-in via
`RERANK_ENABLED`.
Dense P@1 miss: `create a user account` (MRR 0.250).
+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 = dense_pages or bm25_pages # HVM default retrieval is dense-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()