Files
SuperBizAgent-java/scripts/eval_rag_retrieval.py
T

768 lines
27 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. Fixtures must use the
modular `LookupResult` shape produced by lookup_knowledge.
"""
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")
DEFAULT_DIFF_JSON_REPORT = Path("eval/rag-retrieval/reports/baseline-diff.json")
DEFAULT_DIFF_MD_REPORT = Path("eval/rag-retrieval/reports/baseline-diff.md")
HIT_LEVEL_RANK = {
"miss": 0,
"weak": 1,
"medium": 2,
"strong": 3,
}
@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 raw.get("finalRank") or fallback_rank),
doc_id=str(raw.get("docId") or raw.get("source") or raw.get("id") or ""),
title=str(raw.get("title") or ""),
breadcrumb=str(raw.get("breadcrumb") or ""),
content=str(raw.get("content") or raw.get("contentPreview") or ""),
score=_optional_float(
raw.get("score")
if raw.get("score") is not None
else raw.get("finalScore")
),
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}"
@dataclass
class NormalizedFixture:
data_shape: str
candidates: list[Candidate]
selected_attempt: str | None
fallback_reason: str | None
evidence_status: str | None
included_sources: list[str]
omitted_sources: list[str]
rerank_top_source: str | None
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 normalize_sources(values: list[Any]) -> list[str]:
return [str(value) for value in values if str(value).strip()]
def get_lookup_result(fixture: dict[str, Any]) -> dict[str, Any] | None:
lookup = fixture.get("lookupResult")
if isinstance(lookup, dict):
return lookup
if "evidenceBlocks" in fixture or "retrievalTrace" in fixture:
return fixture
return None
def normalize_fixture(fixture: dict[str, Any], top_k: int) -> NormalizedFixture:
lookup_result = get_lookup_result(fixture)
if lookup_result is None:
return NormalizedFixture(
data_shape="invalid",
candidates=[],
selected_attempt=None,
fallback_reason=None,
evidence_status=None,
included_sources=[],
omitted_sources=[],
rerank_top_source=None,
)
raw_blocks = lookup_result.get("evidenceBlocks") or []
candidates = [
Candidate.from_json(raw, index + 1)
for index, raw in enumerate(raw_blocks[:top_k])
if isinstance(raw, dict)
]
retrieval_trace = lookup_result.get("retrievalTrace") or {}
context_pack = lookup_result.get("contextPack") or {}
rerank_trace = lookup_result.get("rerankTrace") or {}
rerank_items = [
item for item in rerank_trace.get("items", [])
if isinstance(item, dict)
]
rerank_items.sort(key=lambda item: int(item.get("finalRank") or 999999))
return NormalizedFixture(
data_shape="lookupResult",
candidates=candidates,
selected_attempt=optional_string(retrieval_trace.get("selectedAttempt")),
fallback_reason=optional_string(retrieval_trace.get("fallbackReason")),
evidence_status=optional_string(retrieval_trace.get("evidenceStatus")),
included_sources=normalize_sources(context_pack.get("includedSources") or []),
omitted_sources=normalize_sources(context_pack.get("omittedSources") or []),
rerank_top_source=(
optional_string(rerank_items[0].get("source"))
if rerank_items
else None
),
)
def optional_string(value: Any) -> str | None:
if value is None:
return None
text = str(value)
return text if text else None
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_sources = normalize_terms(case.get("expectedSources", []))
expected_documents = expected_doc_ids or expected_sources
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"),
"dataShape": None,
"hitLevel": "miss",
"passed": False,
"firstExpectedRank": None,
"topCandidates": [],
"matchedKeywords": [],
"breadcrumbMatched": False,
"selectedAttempt": None,
"fallbackReason": None,
"evidenceStatus": None,
"includedSources": [],
"omittedSources": [],
"rerankTopSource": None,
"failedChecks": [f"missing fixture: {fixture_path.as_posix()}"],
}
fixture = normalize_fixture(load_json(fixture_path), top_k)
candidates = fixture.candidates
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_documents):
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")
failed_checks.extend(check_modular_contract(case, fixture))
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"),
"dataShape": fixture.data_shape,
"hitLevel": hit_level,
"passed": hit_level in {"strong", "medium"} and not failed_checks,
"firstExpectedRank": first_expected,
"topCandidates": [candidate.label() for candidate in candidates],
"matchedKeywords": keyword_matches,
"breadcrumbMatched": breadcrumb_match,
"selectedAttempt": fixture.selected_attempt,
"fallbackReason": fixture.fallback_reason,
"evidenceStatus": fixture.evidence_status,
"includedSources": fixture.included_sources,
"omittedSources": fixture.omitted_sources,
"rerankTopSource": fixture.rerank_top_source,
"failedChecks": failed_checks,
}
def check_modular_contract(case: dict[str, Any], fixture: NormalizedFixture) -> list[str]:
failed: list[str] = []
if fixture.data_shape != "lookupResult":
failed.append("fixture must use lookupResult shape")
compare_expected(
failed,
case,
"expectedSelectedAttempt",
fixture.selected_attempt,
"selected attempt mismatch",
)
if "expectedFallbackReasons" in case:
compare_expected_any(
failed,
case,
"expectedFallbackReasons",
fixture.fallback_reason,
"fallback reason mismatch",
)
else:
compare_expected(
failed,
case,
"expectedFallbackReason",
fixture.fallback_reason,
"fallback reason mismatch",
)
compare_expected(
failed,
case,
"expectedEvidenceStatus",
fixture.evidence_status,
"evidence status mismatch",
)
compare_expected(
failed,
case,
"expectedRerankTopSource",
fixture.rerank_top_source,
"rerank top source mismatch",
)
expected_context_sources = normalize_sources(case.get("expectedContextSources", []))
if expected_context_sources:
included = set(fixture.included_sources)
missing = [
source for source in expected_context_sources
if source not in included
]
if missing:
failed.append("expected context sources missing: " + ", ".join(missing))
return failed
def compare_expected(
failed: list[str],
case: dict[str, Any],
field: str,
actual: str | None,
message: str,
) -> None:
if field not in case:
return
expected = case.get(field)
if expected is None:
if actual is not None:
failed.append(f"{message}: expected <none>, got {actual}")
return
if str(expected) != str(actual):
failed.append(f"{message}: expected {expected}, got {actual or '<none>'}")
def compare_expected_any(
failed: list[str],
case: dict[str, Any],
field: str,
actual: str | None,
message: str,
) -> None:
expected_values = case.get(field)
if not isinstance(expected_values, list):
failed.append(f"{field} must be a list")
return
normalized_expected = [
None if value is None else str(value)
for value in expected_values
]
normalized_actual = None if actual is None else str(actual)
if normalized_actual not in normalized_expected:
expected_text = ", ".join(
"<none>" if value is None else value
for value in normalized_expected
)
failed.append(f"{message}: expected one of [{expected_text}], got {actual or '<none>'}")
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
]
retrieved = counts["strong"] + counts["medium"]
passed = sum(1 for item in results if item["passed"])
lookup_result_cases = sum(
1 for item in results if item.get("dataShape") == "lookupResult"
)
return {
"caseCount": total,
"topK": top_k,
"passedCount": passed,
"failedCount": total - passed,
"passRate": round(passed / total, 4) if total else 0,
"lookupResultCaseCount": lookup_result_cases,
"strongHitCount": counts["strong"],
"mediumHitCount": counts["medium"],
"weakHitCount": counts["weak"],
"missCount": counts["miss"],
"recallAtK": round(retrieved / 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"| Passed | {metrics['passedCount']} |",
f"| Failed | {metrics['failedCount']} |",
f"| Pass rate | {metrics['passRate']} |",
f"| LookupResult fixtures | {metrics['lookupResultCaseCount']} |",
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 | Pass | Hit | Attempt | Fallback | Evidence | 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} | {passed} | {hit} | {attempt} | {fallback} | {evidence} | {rank} | {top} | {failed} |".format(
case=item["caseId"],
scenario=item.get("scenario") or "",
passed=str(item["passed"]).lower(),
hit=item["hitLevel"],
attempt=item.get("selectedAttempt") or "",
fallback=item.get("fallbackReason") or "",
evidence=item.get("evidenceStatus") or "",
rank=first_rank if first_rank is not None else "",
top=top,
failed=failed,
)
)
lines.append("")
return "\n".join(lines)
def compare_reports(baseline: dict[str, Any], current: dict[str, Any]) -> dict[str, Any]:
items: list[dict[str, Any]] = []
compare_metric(items, "aggregate", None, "passRate", baseline, current, higher_is_better=True)
compare_metric(items, "aggregate", None, "recallAtK", baseline, current, higher_is_better=True)
compare_metric(items, "aggregate", None, "strongHitRate", baseline, current, higher_is_better=True)
compare_metric(items, "aggregate", None, "missCount", baseline, current, higher_is_better=False)
compare_cases(items, baseline.get("results") or [], current.get("results") or [])
regression_count = count_type(items, "REGRESSION")
improvement_count = count_type(items, "IMPROVEMENT")
changed_count = count_type(items, "CHANGED")
return {
"generatedAt": datetime.now(timezone.utc).isoformat(),
"baselineReport": baseline.get("caseFile"),
"currentReport": current.get("caseFile"),
"baselineCaseCount": get_aggregate_value(baseline, "caseCount"),
"currentCaseCount": get_aggregate_value(current, "caseCount"),
"baselinePassRate": get_aggregate_value(baseline, "passRate"),
"currentPassRate": get_aggregate_value(current, "passRate"),
"baselineRecallAtK": get_aggregate_value(baseline, "recallAtK"),
"currentRecallAtK": get_aggregate_value(current, "recallAtK"),
"regressionCount": regression_count,
"improvementCount": improvement_count,
"changedCount": changed_count,
"hasRegression": regression_count > 0,
"items": items,
}
def compare_metric(
items: list[dict[str, Any]],
scope: str,
case_id: str | None,
metric: str,
baseline: dict[str, Any],
current: dict[str, Any],
higher_is_better: bool,
) -> None:
baseline_value = get_aggregate_value(baseline, metric)
current_value = get_aggregate_value(current, metric)
if baseline_value == current_value:
return
if not isinstance(baseline_value, (int, float)) or not isinstance(current_value, (int, float)):
items.append(diff_item("CHANGED", scope, case_id, metric, baseline_value, current_value, None))
return
delta = float(current_value) - float(baseline_value)
items.append(diff_item(classify_delta(delta, higher_is_better), scope, case_id, metric,
baseline_value, current_value, delta))
def compare_cases(
items: list[dict[str, Any]],
baseline_results: list[dict[str, Any]],
current_results: list[dict[str, Any]],
) -> None:
baseline_by_id = by_case_id(baseline_results)
current_by_id = by_case_id(current_results)
case_ids = sorted(set(baseline_by_id) | set(current_by_id))
for case_id in case_ids:
baseline = baseline_by_id.get(case_id)
current = current_by_id.get(case_id)
if baseline is None:
items.append(diff_item("CHANGED", "case", case_id, "casePresence", "missing", "present", None))
continue
if current is None:
items.append(diff_item("REGRESSION", "case", case_id, "casePresence", "present", "missing", None))
continue
compare_case_bool(items, case_id, "passed", baseline, current, higher_is_better=True)
compare_hit_level(items, case_id, baseline, current)
compare_case_rank(items, case_id, baseline, current)
compare_case_value(items, case_id, "selectedAttempt", baseline, current)
compare_case_value(items, case_id, "fallbackReason", baseline, current)
compare_case_value(items, case_id, "evidenceStatus", baseline, current)
compare_case_value(items, case_id, "rerankTopSource", baseline, current)
def compare_case_bool(
items: list[dict[str, Any]],
case_id: str,
metric: str,
baseline: dict[str, Any],
current: dict[str, Any],
higher_is_better: bool,
) -> None:
baseline_value = bool(baseline.get(metric))
current_value = bool(current.get(metric))
if baseline_value == current_value:
return
delta = int(current_value) - int(baseline_value)
items.append(diff_item(classify_delta(delta, higher_is_better), "case", case_id, metric,
baseline_value, current_value, float(delta)))
def compare_hit_level(
items: list[dict[str, Any]],
case_id: str,
baseline: dict[str, Any],
current: dict[str, Any],
) -> None:
baseline_value = baseline.get("hitLevel")
current_value = current.get("hitLevel")
if baseline_value == current_value:
return
delta = HIT_LEVEL_RANK.get(str(current_value), 0) - HIT_LEVEL_RANK.get(str(baseline_value), 0)
items.append(diff_item(classify_delta(delta, True), "case", case_id, "hitLevel",
baseline_value, current_value, float(delta)))
def compare_case_rank(
items: list[dict[str, Any]],
case_id: str,
baseline: dict[str, Any],
current: dict[str, Any],
) -> None:
baseline_value = baseline.get("firstExpectedRank")
current_value = current.get("firstExpectedRank")
if baseline_value == current_value:
return
if baseline_value is None or current_value is None:
change_type = "REGRESSION" if current_value is None else "IMPROVEMENT"
items.append(diff_item(change_type, "case", case_id, "firstExpectedRank",
baseline_value, current_value, None))
return
delta = int(current_value) - int(baseline_value)
items.append(diff_item(classify_delta(delta, False), "case", case_id, "firstExpectedRank",
baseline_value, current_value, float(delta)))
def compare_case_value(
items: list[dict[str, Any]],
case_id: str,
metric: str,
baseline: dict[str, Any],
current: dict[str, Any],
) -> None:
baseline_value = baseline.get(metric)
current_value = current.get(metric)
if baseline_value == current_value:
return
items.append(diff_item("CHANGED", "case", case_id, metric, baseline_value, current_value, None))
def diff_item(
change_type: str,
scope: str,
case_id: str | None,
metric: str,
baseline_value: Any,
current_value: Any,
delta: float | None,
) -> dict[str, Any]:
target = case_id or scope
return {
"type": change_type,
"scope": scope,
"caseId": case_id,
"metric": metric,
"baselineValue": value_label(baseline_value),
"currentValue": value_label(current_value),
"delta": delta,
"message": f"{target} {metric} changed",
}
def render_diff_markdown(report: dict[str, Any]) -> str:
lines = [
"# RAG Retrieval Baseline Diff",
"",
f"Generated at: `{report['generatedAt']}`",
"",
"## Summary",
"",
"| Metric | Value |",
"|---|---:|",
f"| Baseline cases | {report['baselineCaseCount']} |",
f"| Current cases | {report['currentCaseCount']} |",
f"| Baseline pass rate | {report['baselinePassRate']} |",
f"| Current pass rate | {report['currentPassRate']} |",
f"| Baseline recall@K | {report['baselineRecallAtK']} |",
f"| Current recall@K | {report['currentRecallAtK']} |",
f"| Regressions | {report['regressionCount']} |",
f"| Improvements | {report['improvementCount']} |",
f"| Changed | {report['changedCount']} |",
"",
"## Items",
"",
"| Type | Scope | Case | Metric | Baseline | Current | Delta | Message |",
"|---|---|---|---|---|---|---:|---|",
]
for item in report["items"]:
delta = "" if item.get("delta") is None else item["delta"]
lines.append(
"| {type} | {scope} | {case} | {metric} | {baseline} | {current} | {delta} | {message} |".format(
type=item["type"],
scope=item["scope"],
case=item.get("caseId") or "",
metric=item["metric"],
baseline=item["baselineValue"],
current=item["currentValue"],
delta=delta,
message=item["message"],
)
)
lines.append("")
return "\n".join(lines)
def get_aggregate_value(report: dict[str, Any], metric: str) -> Any:
return (report.get("aggregate") or {}).get(metric)
def by_case_id(results: list[dict[str, Any]]) -> dict[str, dict[str, Any]]:
return {
str(item.get("caseId")): item
for item in sorted(results, key=lambda item: str(item.get("caseId")))
}
def classify_delta(delta: float, higher_is_better: bool) -> str:
if delta == 0.0:
return "CHANGED"
improved = delta > 0 if higher_is_better else delta < 0
return "IMPROVEMENT" if improved else "REGRESSION"
def count_type(items: list[dict[str, Any]], change_type: str) -> int:
return sum(1 for item in items if item["type"] == change_type)
def value_label(value: Any) -> str:
if value is None:
return "-"
return str(value)
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)
parser.add_argument(
"--compare-to",
type=Path,
default=None,
help="Optional baseline report to diff against the newly generated report.",
)
parser.add_argument("--diff-json-report", type=Path, default=DEFAULT_DIFF_JSON_REPORT)
parser.add_argument("--diff-markdown-report", type=Path, default=DEFAULT_DIFF_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 not item["passed"]]
diff_failed = False
if args.compare_to is not None:
baseline = load_json(args.compare_to)
diff = compare_reports(baseline, report)
write_json(args.diff_json_report, diff)
write_text(args.diff_markdown_report, render_diff_markdown(diff))
diff_failed = bool(diff["hasRegression"])
print(
"Diffed against {baseline}: regressions={regressions}, changes={changes}".format(
baseline=args.compare_to.as_posix(),
regressions=diff["regressionCount"],
changes=diff["changedCount"],
)
)
print(
"Evaluated {total} cases: passRate={pass_rate}, recall@{top_k}={recall}, failed={failed}".format(
total=len(results),
pass_rate=report["aggregate"]["passRate"],
top_k=top_k,
recall=report["aggregate"]["recallAtK"],
failed=len(failed),
)
)
return 1 if failed or diff_failed else 0
if __name__ == "__main__":
raise SystemExit(main())