test: add rag retrieval baseline
This commit is contained in:
@@ -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())
|
||||
Reference in New Issue
Block a user