test: add rag retrieval baseline

This commit is contained in:
aruo
2026-07-05 02:02:27 +08:00
parent 79feed3314
commit 9a2a44d1b5
19 changed files with 946 additions and 0 deletions
+290
View File
@@ -0,0 +1,290 @@
#!/usr/bin/env python3
"""Offline evaluator for RAG retrieval golden cases.
The evaluator reads fixed golden cases and saved retrieval fixtures. It does not
call the running application or any external service.
"""
from __future__ import annotations
import argparse
import json
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
DEFAULT_CASES = Path("eval/rag-retrieval/cases/golden-cases.json")
DEFAULT_FIXTURES = Path("eval/rag-retrieval/fixtures")
DEFAULT_JSON_REPORT = Path("eval/rag-retrieval/reports/baseline.json")
DEFAULT_MD_REPORT = Path("eval/rag-retrieval/reports/baseline.md")
@dataclass
class Candidate:
rank: int
doc_id: str
title: str
breadcrumb: str
content: str
score: float | None
retrieval_layer: str | None
@classmethod
def from_json(cls, raw: dict[str, Any], fallback_rank: int) -> "Candidate":
return cls(
rank=int(raw.get("rank") or fallback_rank),
doc_id=str(raw.get("docId") or raw.get("id") or ""),
title=str(raw.get("title") or ""),
breadcrumb=str(raw.get("breadcrumb") or ""),
content=str(raw.get("content") or ""),
score=_optional_float(raw.get("score")),
retrieval_layer=(
str(raw.get("retrievalLayer"))
if raw.get("retrievalLayer") is not None
else None
),
)
def searchable_text(self) -> str:
return " ".join(
[self.doc_id, self.title, self.breadcrumb, self.content]
).lower()
def label(self) -> str:
label = self.doc_id or self.title or f"rank-{self.rank}"
return f"{self.rank}:{label}"
def _optional_float(value: Any) -> float | None:
if value is None:
return None
try:
return float(value)
except (TypeError, ValueError):
return None
def load_json(path: Path) -> Any:
with path.open("r", encoding="utf-8") as handle:
return json.load(handle)
def write_json(path: Path, payload: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8", newline="\n") as handle:
json.dump(payload, handle, ensure_ascii=False, indent=2)
handle.write("\n")
def write_text(path: Path, content: str) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8", newline="\n") as handle:
handle.write(content)
def normalize_terms(values: list[Any]) -> list[str]:
return [str(value).lower() for value in values if str(value).strip()]
def evaluate_case(case: dict[str, Any], fixture_dir: Path, top_k: int) -> dict[str, Any]:
case_id = str(case["caseId"])
fixture_path = fixture_dir / f"{case_id}.json"
expected_doc_ids = normalize_terms(case.get("expectedDocIds", []))
expected_breadcrumbs = normalize_terms(case.get("expectedBreadcrumbs", []))
expected_keywords = normalize_terms(case.get("expectedKeywords", []))
if not fixture_path.exists():
return {
"caseId": case_id,
"scenario": case.get("scenario"),
"query": case.get("query"),
"hitLevel": "miss",
"passed": False,
"firstExpectedRank": None,
"topCandidates": [],
"failedChecks": [f"missing fixture: {fixture_path.as_posix()}"],
}
fixture = load_json(fixture_path)
raw_candidates = fixture.get("candidates", [])
candidates = [
Candidate.from_json(raw, index + 1)
for index, raw in enumerate(raw_candidates[:top_k])
]
first_expected = None
expected_doc_candidate = None
for candidate in candidates:
candidate_doc = candidate.doc_id.lower()
if any(expected == candidate_doc for expected in expected_doc_ids):
first_expected = candidate.rank
expected_doc_candidate = candidate
break
breadcrumb_match = False
keyword_matches: list[str] = []
if expected_doc_candidate is not None:
breadcrumb_text = expected_doc_candidate.breadcrumb.lower()
breadcrumb_match = any(
expected in breadcrumb_text or breadcrumb_text in expected
for expected in expected_breadcrumbs
)
searchable = expected_doc_candidate.searchable_text()
keyword_matches = [
keyword for keyword in expected_keywords if keyword in searchable
]
else:
all_text = " ".join(candidate.searchable_text() for candidate in candidates)
keyword_matches = [keyword for keyword in expected_keywords if keyword in all_text]
failed_checks: list[str] = []
if expected_doc_candidate is None:
failed_checks.append("expected document not found")
if expected_doc_candidate is not None and expected_breadcrumbs and not breadcrumb_match:
failed_checks.append("expected breadcrumb not found on expected document")
if expected_keywords and not keyword_matches:
failed_checks.append("expected evidence keywords not found")
if expected_doc_candidate is not None and (
breadcrumb_match or bool(keyword_matches)
):
hit_level = "strong"
elif expected_doc_candidate is not None:
hit_level = "medium"
elif keyword_matches:
hit_level = "weak"
else:
hit_level = "miss"
return {
"caseId": case_id,
"scenario": case.get("scenario"),
"query": case.get("query"),
"hitLevel": hit_level,
"passed": hit_level in {"strong", "medium"},
"firstExpectedRank": first_expected,
"topCandidates": [candidate.label() for candidate in candidates],
"matchedKeywords": keyword_matches,
"breadcrumbMatched": breadcrumb_match,
"failedChecks": failed_checks,
}
def aggregate(results: list[dict[str, Any]], top_k: int) -> dict[str, Any]:
total = len(results)
counts = {
"strong": sum(1 for item in results if item["hitLevel"] == "strong"),
"medium": sum(1 for item in results if item["hitLevel"] == "medium"),
"weak": sum(1 for item in results if item["hitLevel"] == "weak"),
"miss": sum(1 for item in results if item["hitLevel"] == "miss"),
}
expected_ranks = [
item["firstExpectedRank"]
for item in results
if item.get("firstExpectedRank") is not None
]
passed = counts["strong"] + counts["medium"]
return {
"caseCount": total,
"topK": top_k,
"strongHitCount": counts["strong"],
"mediumHitCount": counts["medium"],
"weakHitCount": counts["weak"],
"missCount": counts["miss"],
"recallAtK": round(passed / total, 4) if total else 0,
"strongHitRate": round(counts["strong"] / total, 4) if total else 0,
"averageFirstHitRank": (
round(sum(expected_ranks) / len(expected_ranks), 4)
if expected_ranks
else None
),
}
def render_markdown(report: dict[str, Any]) -> str:
metrics = report["aggregate"]
lines = [
"# RAG Retrieval Baseline",
"",
f"Generated at: `{report['generatedAt']}`",
"",
"## Aggregate",
"",
"| Metric | Value |",
"|---|---:|",
f"| Cases | {metrics['caseCount']} |",
f"| Top K | {metrics['topK']} |",
f"| Recall@K | {metrics['recallAtK']} |",
f"| Strong hit rate | {metrics['strongHitRate']} |",
f"| Strong hits | {metrics['strongHitCount']} |",
f"| Medium hits | {metrics['mediumHitCount']} |",
f"| Weak hits | {metrics['weakHitCount']} |",
f"| Misses | {metrics['missCount']} |",
f"| Average first hit rank | {metrics['averageFirstHitRank']} |",
"",
"## Cases",
"",
"| Case | Scenario | Hit | First Expected Rank | Top Candidates | Failed Checks |",
"|---|---|---|---:|---|---|",
]
for item in report["results"]:
failed = "<br>".join(item["failedChecks"]) if item["failedChecks"] else ""
top = "<br>".join(item["topCandidates"])
first_rank = item["firstExpectedRank"]
lines.append(
"| {case} | {scenario} | {hit} | {rank} | {top} | {failed} |".format(
case=item["caseId"],
scenario=item.get("scenario") or "",
hit=item["hitLevel"],
rank=first_rank if first_rank is not None else "",
top=top,
failed=failed,
)
)
lines.append("")
return "\n".join(lines)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--cases", type=Path, default=DEFAULT_CASES)
parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES)
parser.add_argument("--json-report", type=Path, default=DEFAULT_JSON_REPORT)
parser.add_argument("--markdown-report", type=Path, default=DEFAULT_MD_REPORT)
return parser.parse_args()
def main() -> int:
args = parse_args()
case_file = load_json(args.cases)
cases = case_file.get("cases", [])
top_k = int(case_file.get("topK") or 5)
results = [evaluate_case(case, args.fixtures, top_k) for case in cases]
report = {
"generatedAt": datetime.now(timezone.utc).isoformat(),
"caseFile": args.cases.as_posix(),
"fixtureDir": args.fixtures.as_posix(),
"aggregate": aggregate(results, top_k),
"results": results,
}
write_json(args.json_report, report)
write_text(args.markdown_report, render_markdown(report))
failed = [item for item in results if item["hitLevel"] == "miss"]
print(
"Evaluated {total} cases: recall@{top_k}={recall}, misses={misses}".format(
total=len(results),
top_k=top_k,
recall=report["aggregate"]["recallAtK"],
misses=len(failed),
)
)
return 1 if failed else 0
if __name__ == "__main__":
raise SystemExit(main())