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