feat(rag): close eval pipeline with live snapshots

This commit is contained in:
zhuyongxin
2026-07-06 21:39:27 +08:00
parent cf3333d607
commit ed7efc58b7
47 changed files with 2613 additions and 177 deletions
+499 -22
View File
@@ -2,7 +2,8 @@
"""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.
call the running application or any external service. Fixtures must use the
modular `LookupResult` shape produced by lookup_knowledge.
"""
from __future__ import annotations
@@ -19,6 +20,15 @@ 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
@@ -34,12 +44,16 @@ class Candidate:
@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 ""),
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 ""),
score=_optional_float(raw.get("score")),
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
@@ -57,6 +71,18 @@ class Candidate:
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
@@ -88,10 +114,76 @@ 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", []))
@@ -100,25 +192,30 @@ def evaluate_case(case: dict[str, Any], fixture_dir: Path, top_k: int) -> dict[s
"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 = 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])
]
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_doc_ids):
if any(expected == candidate_doc for expected in expected_documents):
first_expected = candidate.rank
expected_doc_candidate = candidate
break
@@ -149,6 +246,8 @@ def evaluate_case(case: dict[str, Any], fixture_dir: Path, top_k: int) -> dict[s
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)
):
@@ -164,16 +263,119 @@ def evaluate_case(case: dict[str, Any], fixture_dir: Path, top_k: int) -> dict[s
"caseId": case_id,
"scenario": case.get("scenario"),
"query": case.get("query"),
"dataShape": fixture.data_shape,
"hitLevel": hit_level,
"passed": hit_level in {"strong", "medium"},
"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 = {
@@ -187,15 +389,23 @@ def aggregate(results: list[dict[str, Any]], top_k: int) -> dict[str, Any]:
for item in results
if item.get("firstExpectedRank") is not None
]
passed = counts["strong"] + counts["medium"]
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(passed / total, 4) if total else 0,
"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)
@@ -218,6 +428,10 @@ def render_markdown(report: dict[str, Any]) -> str:
"|---|---:|",
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']} |",
@@ -228,18 +442,22 @@ def render_markdown(report: dict[str, Any]) -> str:
"",
"## Cases",
"",
"| Case | Scenario | Hit | First Expected Rank | Top Candidates | Failed Checks |",
"|---|---|---|---:|---|---|",
"| 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} | {hit} | {rank} | {top} | {failed} |".format(
"| {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,
@@ -249,12 +467,255 @@ def render_markdown(report: dict[str, Any]) -> str:
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()
@@ -274,16 +735,32 @@ def main() -> int:
write_json(args.json_report, report)
write_text(args.markdown_report, render_markdown(report))
failed = [item for item in results if item["hitLevel"] == "miss"]
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: recall@{top_k}={recall}, misses={misses}".format(
"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"],
misses=len(failed),
failed=len(failed),
)
)
return 1 if failed else 0
return 1 if failed or diff_failed else 0
if __name__ == "__main__":