0.1.1 暂存

This commit is contained in:
陈辅元
2026-04-14 10:28:22 +08:00
parent 4a0638ba2e
commit cc83fe963a
83 changed files with 4005 additions and 179 deletions
+364
View File
@@ -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 {}