"""NL2SQL 离线 Golden Case 评测。""" from __future__ import annotations import json from pathlib import Path from sqlglot import parse_one from nl2sql.sql_security import SqlSecurityError, validate_select_sql def load_cases(path: str | Path) -> list[dict]: """读取结构化评测案例,不读取真实查询结果。""" payload = json.loads(Path(path).read_text(encoding="utf-8")) if not isinstance(payload, list): raise ValueError("Golden Case 必须是数组") return payload def evaluate_case(case: dict) -> dict: """评估单个案例的表召回和 SQL 安全结果。""" expected = set(case.get("expected_tables", [])) retrieved = set(case.get("retrieved_tables", [])) recall = 1.0 if not expected else len(expected & retrieved) / len(expected) result = { "case_id": str(case.get("case_id", "unknown")), "table_recall": recall, "security_pass": False, } try: validated = validate_select_sql( case.get("sql", ""), authorized_tables=set(case.get("authorized_tables", [])), authorized_columns=case.get("authorized_columns"), max_rows=int(case.get("max_rows", 1000)), ) result["security_pass"] = True result["access_tables"] = sorted(validated.access_tables) except SqlSecurityError as exc: result["error_type"] = type(exc).__name__ except (TypeError, ValueError, KeyError) as exc: result["error_type"] = type(exc).__name__ if "expected_sql" in case and "sql" in case: result["sql_match"] = compare_sql(case["sql"], case["expected_sql"]) if "expected_result" in case and "actual_result" in case: result["result_match"] = compare_result_set(case["expected_result"], case["actual_result"]) return result def evaluate_cases(cases: list[dict]) -> list[dict]: return [evaluate_case(case) for case in cases] def summarize_evaluation(results: list[dict]) -> dict: total = len(results) passed = sum(1 for item in results if item.get("security_pass")) recall = sum(float(item.get("table_recall", 0.0)) for item in results) summary = { "total": total, "security_passed": passed, "security_pass_rate": passed / total if total else 0.0, "average_table_recall": recall / total if total else 0.0, "failed_case_ids": [item["case_id"] for item in results if not item.get("security_pass")], } sql_results = [item["sql_match"] for item in results if "sql_match" in item] result_results = [item["result_match"] for item in results if "result_match" in item] if sql_results: summary["sql_match_rate"] = sum(sql_results) / len(sql_results) if result_results: summary["result_match_rate"] = sum(result_results) / len(result_results) return summary def compare_sql(actual_sql: str, expected_sql: str) -> bool: """使用 AST 规范化后比较 SQL,不连接数据库。""" try: return parse_one(actual_sql, read="mysql").sql(dialect="mysql") == parse_one( expected_sql, read="mysql" ).sql(dialect="mysql") except (TypeError, ValueError): return False def compare_result_set(expected: dict, actual: dict) -> bool: """比较假结果集的列名和行数据,不执行 SQL。""" if not isinstance(expected, dict) or not isinstance(actual, dict): return False return ( expected.get("columns", []) == actual.get("columns", []) and expected.get("rows", []) == actual.get("rows", []) ) def build_evaluation_report( cases: list[dict], *, prompt_version: str = "unknown", semantic_version: str = "unknown", model_version: str = "unknown", ) -> dict: """构建带版本元数据的离线评测报告。""" results = evaluate_cases(cases) return { "metadata": { "prompt_version": prompt_version, "semantic_version": semantic_version, "model_version": model_version, }, "summary": summarize_evaluation(results), "cases": results, }