Files
ai-g3sb-backman2.0/utils/sql_parser.py
T
2026-04-10 16:52:07 +08:00

365 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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 {}