0.1.1 暂存
This commit is contained in:
@@ -0,0 +1,364 @@
|
||||
"""
|
||||
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 {}
|
||||
Reference in New Issue
Block a user