admesh/agentic-intent-classifier
254
1from __future__ import annotations2 3import json4import sys5from pathlib import Path6 7BASE_DIR = Path(__file__).resolve().parent.parent8if str(BASE_DIR) not in sys.path:9 sys.path.insert(0, str(BASE_DIR))10 11from combined_inference import classify_query12from schemas import validate_classify_response13 14 15def load_cases(path: Path) -> list[dict]:16 return json.loads(path.read_text(encoding="utf-8"))17 18 19def write_json(path: Path, payload: dict | list) -> None:20 path.parent.mkdir(parents=True, exist_ok=True)21 path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8")22 23 24def resolve_path(payload: dict, dotted_path: str):25 value = payload26 for part in dotted_path.split("."):27 if isinstance(value, dict):28 value = value.get(part)29 else:30 return None31 return value32 33 34def evaluate_case_file(cases_path: Path, output_dir: Path, artifact_name: str) -> dict:35 cases = load_cases(cases_path)36 results = []37 counts_by_status: dict[str, dict[str, int]] = {}38 39 for case in cases:40 payload = validate_classify_response(classify_query(case["text"]))41 mismatches = []42 expected = case.get("expected", {})43 actual_snapshot = {}44 for dotted_path, expected_value in expected.items():45 actual_value = resolve_path(payload, dotted_path)46 actual_snapshot[dotted_path] = actual_value47 if actual_value != expected_value:48 mismatches.append(49 {50 "path": dotted_path,51 "expected": expected_value,52 "actual": actual_value,53 }54 )55 56 status = case["status"]57 bucket = counts_by_status.setdefault(status, {"total": 0, "passed": 0, "failed": 0})58 bucket["total"] += 159 if mismatches:60 bucket["failed"] += 161 else:62 bucket["passed"] += 163 64 results.append(65 {66 "id": case["id"],67 "status": status,68 "text": case["text"],69 "notes": case.get("notes", ""),70 "pass": not mismatches,71 "mismatches": mismatches,72 "expected": expected,73 "actual": actual_snapshot,74 }75 )76 77 summary = {78 "cases_path": str(cases_path),79 "count": len(results),80 "passed": sum(1 for item in results if item["pass"]),81 "failed": sum(1 for item in results if not item["pass"]),82 "by_status": counts_by_status,83 "results": results,84 }85 write_json(output_dir / artifact_name, summary)86 return summary87 88 89def evaluate_known_failure_cases(cases_path: Path, output_dir: Path) -> dict:90 return evaluate_case_file(cases_path, output_dir, "known_failure_regression.json")91 92 93def evaluate_iab_behavior_lock_cases(cases_path: Path, output_dir: Path) -> dict:94 return evaluate_case_file(cases_path, output_dir, "iab_behavior_lock_regression.json")95 96 97def evaluate_iab_cross_vertical_behavior_lock_cases(cases_path: Path, output_dir: Path) -> dict:98 return evaluate_case_file(cases_path, output_dir, "iab_cross_vertical_behavior_lock_regression.json")99 100 101def evaluate_iab_quality_target_cases(cases_path: Path, output_dir: Path) -> dict:102 return evaluate_case_file(cases_path, output_dir, "iab_quality_target_eval.json")103 104 105def evaluate_iab_cross_vertical_quality_target_cases(cases_path: Path, output_dir: Path) -> dict:106 return evaluate_case_file(cases_path, output_dir, "iab_cross_vertical_quality_target_eval.json")107 