365 lines
10 KiB
Python
365 lines
10 KiB
Python
"""
|
||
SQL 解析与验证工具(基于 sqlglot)
|
||
"""
|
||
|
||
import logging
|
||
import re
|
||
from typing import List, Tuple, Optional, Dict
|
||
import sqlglot
|
||
from sqlglot import exp, parse_one, ParseError
|
||
|
||
from schema.manager import SchemaManager
|
||
from schema.models import Table, Column
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
def validate_sql_syntax(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]:
|
||
"""
|
||
验证SQL语法是否正确
|
||
|
||
Args:
|
||
sql: SQL语句
|
||
dialect: SQL方言
|
||
|
||
Returns:
|
||
(是否有效, 错误信息列表)
|
||
"""
|
||
errors = []
|
||
|
||
try:
|
||
# 尝试解析
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
|
||
if parsed is None:
|
||
errors.append("SQL解析返回空结果")
|
||
return False, errors
|
||
|
||
# 检查是否为只读查询(SELECT/CTE/SHOW等)
|
||
# 动态获取可用表达式类型(兼容不同sqlglot版本)
|
||
readable_ops = [exp.Select, exp.Union, exp.Intersect, exp.Except, exp.With]
|
||
# 可选:添加 Show, Describe, Explain(如果存在)
|
||
for op_name in ['Show', 'Describe', 'Explain']:
|
||
if hasattr(exp, op_name):
|
||
readable_ops.append(getattr(exp, op_name))
|
||
|
||
if not isinstance(parsed, tuple(readable_ops)):
|
||
op_type = type(parsed).__name__
|
||
errors.append(f"非查询操作({op_type}),只允许SELECT等只读语句")
|
||
|
||
return True, []
|
||
|
||
except ParseError as e:
|
||
errors.append(f"SQL语法错误: {str(e)}")
|
||
return False, errors
|
||
except Exception as e:
|
||
errors.append(f"解析异常: {str(e)}")
|
||
return False, errors
|
||
|
||
|
||
def extract_tables_from_sql(sql: str, dialect: str = "tsql") -> List[str]:
|
||
"""
|
||
从SQL中提取所有表名
|
||
|
||
Args:
|
||
sql: SQL语句
|
||
dialect: SQL方言
|
||
|
||
Returns:
|
||
表名列表(去重)
|
||
"""
|
||
try:
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
tables = []
|
||
|
||
# 遍历AST查找所有表名
|
||
for node in parsed.walk():
|
||
if isinstance(node, exp.Table):
|
||
table_name = node.name
|
||
if table_name and table_name not in tables:
|
||
tables.append(table_name)
|
||
|
||
return tables
|
||
except Exception as e:
|
||
logger.warning(f"提取表名失败: {e}")
|
||
return []
|
||
|
||
|
||
def build_table_alias_map(parsed: exp.Expression) -> Dict[str, str]:
|
||
"""
|
||
从已解析的 AST 构建「别名/表名 -> 物理表名」映射。
|
||
|
||
FROM T a 时 a -> T,且 T -> T,便于将 a.col 解析到表 T 的列。
|
||
"""
|
||
alias_map: Dict[str, str] = {}
|
||
for node in parsed.walk():
|
||
if not isinstance(node, exp.Table):
|
||
continue
|
||
physical = node.name
|
||
if not physical:
|
||
continue
|
||
alias_map[physical] = physical
|
||
talias = node.args.get("alias")
|
||
if talias is not None:
|
||
aname = talias.name
|
||
if aname:
|
||
alias_map[aname] = physical
|
||
return alias_map
|
||
|
||
|
||
def extract_columns_from_sql(sql: str, dialect: str = "tsql") -> List[Tuple[str, str]]:
|
||
"""
|
||
从SQL中提取所有字段引用(表.字段)
|
||
|
||
Args:
|
||
sql: SQL语句
|
||
dialect: SQL方言
|
||
|
||
Returns:
|
||
[(表名, 字段名), ...] 列表
|
||
"""
|
||
columns = []
|
||
|
||
try:
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
|
||
for node in parsed.walk():
|
||
if isinstance(node, exp.Column):
|
||
table_name = node.table
|
||
col_name = node.name
|
||
if table_name and col_name:
|
||
columns.append((table_name, col_name))
|
||
|
||
return columns
|
||
except Exception as e:
|
||
logger.warning(f"提取字段失败: {e}")
|
||
return []
|
||
|
||
|
||
def validate_schema_consistency(
|
||
sql: str,
|
||
schema_manager: SchemaManager,
|
||
dialect: str = "tsql"
|
||
) -> Tuple[bool, List[str]]:
|
||
"""
|
||
验证SQL与Schema的一致性
|
||
|
||
检查:
|
||
1. 所有表名存在于Schema
|
||
2. 所有字段名属于对应的表
|
||
3. JOIN条件字段存在且类型兼容
|
||
|
||
Args:
|
||
sql: SQL语句
|
||
schema_manager: Schema管理器
|
||
dialect: SQL方言
|
||
|
||
Returns:
|
||
(是否一致, 错误信息列表)
|
||
"""
|
||
errors = []
|
||
|
||
try:
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
except Exception as e:
|
||
logger.debug(f"Schema一致性检查跳过(解析失败): {e}")
|
||
return True, []
|
||
|
||
alias_map = build_table_alias_map(parsed)
|
||
|
||
tables_used: List[str] = []
|
||
for node in parsed.walk():
|
||
if isinstance(node, exp.Table):
|
||
tname = node.name
|
||
if tname and tname not in tables_used:
|
||
tables_used.append(tname)
|
||
|
||
columns_used: List[Tuple[str, str]] = []
|
||
for node in parsed.walk():
|
||
if isinstance(node, exp.Column):
|
||
tref, cname = node.table, node.name
|
||
if tref and cname:
|
||
columns_used.append((tref, cname))
|
||
|
||
# 检查表存在性
|
||
for tbl in tables_used:
|
||
if not schema_manager.get_table(tbl):
|
||
errors.append(f"表不存在: '{tbl}'")
|
||
|
||
# 检查字段存在性(表引用可为物理表名或别名)
|
||
for tbl_name, col_name in columns_used:
|
||
physical = alias_map.get(tbl_name, tbl_name)
|
||
table = schema_manager.get_table(physical)
|
||
if not table:
|
||
errors.append(
|
||
f"字段不存在: '{tbl_name}.{col_name}'"
|
||
f"(无法将表引用解析到已加载Schema中的表)"
|
||
)
|
||
continue
|
||
col_names = [c.name for c in table.columns]
|
||
if col_name not in col_names:
|
||
errors.append(
|
||
f"字段不存在: '{tbl_name}.{col_name}'"
|
||
f"(表 '{physical}' 可用字段: {col_names[:5]}...)"
|
||
)
|
||
|
||
# 检查JOIN条件(外键匹配)
|
||
try:
|
||
for join in parsed.find_all(exp.Join):
|
||
# 解析ON条件
|
||
on_condition = join.args.get("on")
|
||
if on_condition:
|
||
# 检查ON条件中涉及的字段
|
||
for eq in on_condition.find_all(exp.EQ):
|
||
left = eq.left
|
||
right = eq.right
|
||
|
||
# 提取左右两边的表.字段
|
||
for side in [left, right]:
|
||
if isinstance(side, exp.Column):
|
||
tbl = side.table
|
||
col = side.name
|
||
physical = alias_map.get(tbl, tbl)
|
||
table = schema_manager.get_table(physical)
|
||
if table and col not in [c.name for c in table.columns]:
|
||
errors.append(f"JOIN条件字段不存在: {tbl}.{col}")
|
||
except Exception as e:
|
||
logger.debug(f"JOIN条件检查异常: {e}")
|
||
|
||
return len(errors) == 0, errors
|
||
|
||
|
||
def rewrite_mysql_builtins_for_tsql(sql: str) -> str:
|
||
"""
|
||
模型在 T-SQL 目标下仍常输出 MySQL 函数;sqlglot 转写也可能遗漏。
|
||
SQL Server 无 CURDATE()/NOW(),需替换为 GETDATE 族。
|
||
"""
|
||
if not sql:
|
||
return sql
|
||
out = sql
|
||
out = re.sub(
|
||
r"\bCURDATE\s*\(\s*\)",
|
||
"CAST(GETDATE() AS DATE)",
|
||
out,
|
||
flags=re.IGNORECASE,
|
||
)
|
||
out = re.sub(
|
||
r"\bCURRENT_DATE\b",
|
||
"CAST(GETDATE() AS DATE)",
|
||
out,
|
||
flags=re.IGNORECASE,
|
||
)
|
||
out = re.sub(r"\bNOW\s*\(\s*\)", "GETDATE()", out, flags=re.IGNORECASE)
|
||
return out
|
||
|
||
|
||
def normalize_sql_for_dialect(sql: str, dialect: str) -> str:
|
||
"""
|
||
将模型输出的 SQL 规范为目标方言。
|
||
对于 T-SQL,主要进行 MySQL 函数替换(因为模型仍可能输出 CURDATE() 等)。
|
||
"""
|
||
sql = (sql or "").strip()
|
||
if not sql:
|
||
return sql
|
||
|
||
# 如果目标是 T-SQL,只做函数名替换,不再用 sqlglot 转写
|
||
if dialect == "tsql":
|
||
return rewrite_mysql_builtins_for_tsql(sql)
|
||
|
||
return sql
|
||
|
||
|
||
def format_sql(sql: str, dialect: str = "tsql", indent: int = 2) -> str:
|
||
"""
|
||
格式化SQL(可读性)
|
||
|
||
Args:
|
||
sql: SQL语句
|
||
dialect: SQL方言
|
||
indent: 缩进空格数
|
||
|
||
Returns:
|
||
格式化后的SQL
|
||
"""
|
||
try:
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
return parsed.sql(dialect=dialect, pretty=True, indent=indent)
|
||
except Exception as e:
|
||
logger.warning(f"SQL格式化失败: {e}")
|
||
return sql
|
||
|
||
|
||
def normalize_sql(sql: str, dialect: str = "tsql") -> str:
|
||
"""
|
||
标准化SQL(用于比较去重)
|
||
|
||
去除多余空格、统一引号、移除注释等
|
||
|
||
Args:
|
||
sql: SQL语句
|
||
dialect: SQL方言
|
||
|
||
Returns:
|
||
标准化后的SQL
|
||
"""
|
||
try:
|
||
# 解析后重新生成(会规范化格式)
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
normalized = parsed.sql(dialect=dialect, pretty=False)
|
||
# 转换为大写关键词
|
||
return normalized.upper()
|
||
except Exception:
|
||
# 降级:简单处理
|
||
import re
|
||
# 移除多余空格
|
||
sql = re.sub(r'\s+', ' ', sql.strip())
|
||
# 移除注释
|
||
sql = re.sub(r'--.*?$', '', sql, flags=re.MULTILINE)
|
||
sql = re.sub(r'/\*.*?\*/', '', sql, flags=re.DOTALL)
|
||
return sql.upper()
|
||
|
||
|
||
def count_joins(sql: str, dialect: str = "tsql") -> int:
|
||
"""统计JOIN数量"""
|
||
try:
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
joins = list(parsed.find_all(exp.Join))
|
||
return len(joins)
|
||
except Exception:
|
||
return 0
|
||
|
||
|
||
def has_subquery(sql: str, dialect: str = "tsql") -> bool:
|
||
"""检查是否包含子查询"""
|
||
try:
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
# 检查嵌套的SELECT
|
||
for select in parsed.find_all(exp.Select):
|
||
if select is not parsed: # 不是最外层的SELECT
|
||
return True
|
||
return False
|
||
except Exception:
|
||
return False
|
||
|
||
|
||
def get_query_complexity(sql: str, dialect: str = "tsql") -> Dict[str, int]:
|
||
"""
|
||
评估查询复杂度
|
||
|
||
Returns:
|
||
复杂度指标字典
|
||
"""
|
||
try:
|
||
parsed = parse_one(sql, dialect=dialect)
|
||
|
||
return {
|
||
"join_count": len(list(parsed.find_all(exp.Join))),
|
||
"subquery_count": len([s for s in parsed.find_all(exp.Select) if s is not parsed]),
|
||
"where_conditions": len(list(parsed.find_all(exp.Predicate))),
|
||
"aggregation_functions": len(list(parsed.find_all(exp.AggFunc))),
|
||
"column_count": len(list(parsed.find_all(exp.Column))),
|
||
}
|
||
except Exception as e:
|
||
logger.warning(f"复杂度评估失败: {e}")
|
||
return {}
|