291 lines
9.9 KiB
Python
291 lines
9.9 KiB
Python
#!/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())
|