Files
Mutual_Fund/nl2sql/evaluation.py
T

112 lines
4.0 KiB
Python

"""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,
}