""" SQL 验证工具集 """ import re import logging from typing import Tuple, List, Dict logger = logging.getLogger(__name__) # 危险操作关键词(除非明确允许) DANGEROUS_KEYWORDS = [ "DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE", "CREATE", "DROP DATABASE", "DROP TABLE", "DROP INDEX", "GRANT", "REVOKE", "PURGE", "FLUSH", "KILL" ] # 允许的操作(仅查询) ALLOWED_KEYWORDS = [ "SELECT", "WITH", "FROM", "WHERE", "JOIN", "LEFT JOIN", "RIGHT JOIN", "INNER JOIN", "OUTER JOIN", "ON", "USING", "GROUP BY", "HAVING", "ORDER BY", "LIMIT", "OFFSET", "UNION", "UNION ALL", "EXCEPT", "INTERSECT", "AS", "CASE", "WHEN", "THEN", "ELSE", "END", "COUNT", "SUM", "AVG", "MIN", "MAX", "DISTINCT", "AND", "OR", "NOT", "IN", "EXISTS", "BETWEEN", "LIKE", "IS NULL", "IS NOT NULL", "CAST", "COALESCE", "NULLIF", "IFNULL", "DATE", "TIME", "TIMESTAMP", "EXTRACT", "DATE_FORMAT", "STR_TO_DATE", "CURRENT_DATE", "CURRENT_TIMESTAMP", ] def check_dangerous_operations(sql: str) -> Tuple[bool, List[str]]: """ 检查SQL是否包含危险操作 Args: sql: SQL语句(大小写不敏感) Returns: (是否安全, 危险关键词列表) """ sql_upper = sql.upper() found_dangers = [] for keyword in DANGEROUS_KEYWORDS: # 使用正则避免部分匹配(如"DROP"不应匹配"DROPOUT") pattern = r'\b' + re.escape(keyword) + r'\b' if re.search(pattern, sql_upper): found_dangers.append(keyword) is_safe = len(found_dangers) == 0 if not is_safe: logger.warning(f"检测到危险操作: {found_dangers}") return is_safe, found_dangers def validate_no_dml(sql: str) -> Tuple[bool, str]: """ 验证SQL不是DML/DDL操作(仅允许SELECT等查询) Returns: (是否通过, 错误消息) """ sql_upper = sql.strip().upper() # 检查是否以危险关键词开头 first_word = sql_upper.split()[0] if sql_upper.split() else "" if first_word in ["INSERT", "UPDATE", "DELETE", "DROP", "ALTER", "CREATE", "TRUNCATE"]: return False, f"禁止的操作: {first_word}" is_safe, dangers = check_dangerous_operations(sql) if not is_safe: return False, f"SQL包含危险操作: {', '.join(dangers)}" return True, "" def check_sql_injection_patterns(sql: str) -> List[str]: """ 检查明显的SQL注入模式 Args: sql: SQL语句 Returns: 发现的注入模式列表 """ patterns = { "union_all_injection": r"UNION\s+ALL\s+SELECT", "union_injection": r"UNION\s+SELECT", "comment_injection": r"(--|\#|/\*).*SELECT", "semicolon_injection": r";\s*(DROP|DELETE|UPDATE|INSERT)", "or_true_condition": r"OR\s+['\"]?\s*1\s*['\"]?\s*=\s*1", "always_true": r"1\s*=\s*1", } findings = [] sql_lower = sql.lower() for name, pattern in patterns.items(): if re.search(pattern, sql, re.IGNORECASE): findings.append(name) return findings def validate_aggregation_groupby(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]: """ 验证聚合查询的GROUP BY正确性 检查:SELECT中的非聚合字段是否都在GROUP BY中 Args: sql: SQL语句 dialect: SQL方言 Returns: (是否有效, 错误列表) """ from utils.sql_parser import parse_one, exp errors = [] try: parsed = parse_one(sql, dialect=dialect) # 只检查SELECT语句 if not isinstance(parsed, exp.Select): return True, [] # 获取SELECT列表中的表达式 select_exprs = parsed.expressions # 获取GROUP BY字段 group_by = parsed.args.get("group") if not group_by: # 没有GROUP BY但有聚合函数,通常是错误的 has_agg = any( expr.find(exp.AggFunc) is not None for expr in select_exprs ) if has_agg: errors.append("包含聚合函数但缺少GROUP BY子句") return len(errors) == 0, errors group_by_exprs = group_by.expressions # 提取GROUP BY的字段名(简单处理) group_by_cols = set() for expr in group_by_exprs: if isinstance(expr, exp.Column): group_by_cols.add(expr.name) elif isinstance(expr, exp.Ordered): # GROUP BY x ASC/DESC this = expr.this if isinstance(this, exp.Column): group_by_cols.add(this.name) # 检查每个SELECT表达式 for expr in select_exprs: # 如果是聚合函数,跳过 if expr.find(exp.AggFunc): continue # 如果是字面量或表达式,跳过 if isinstance(expr, exp.Literal): continue # 如果是列引用,检查是否在GROUP BY中 if isinstance(expr, exp.Column): col_name = expr.name if col_name not in group_by_cols: errors.append( f"字段 '{col_name}' 在SELECT中但不在GROUP BY中" ) elif isinstance(expr, exp.Alias): # 别名: column AS alias this = expr.this if isinstance(this, exp.Column): col_name = this.name if col_name not in group_by_cols: errors.append( f"字段 '{col_name}' (别名为'{expr.alias}') 在SELECT中但不在GROUP BY中" ) except Exception as e: logger.debug(f"GROUP BY验证异常: {e}") return len(errors) == 0, errors def check_join_conditions(sql: str, dialect: str = "tsql") -> List[str]: """ 检查JOIN条件是否完整 Args: sql: SQL语句 dialect: SQL方言 Returns: 问题列表(空表示无问题) """ from utils.sql_parser import parse_one, exp issues = [] try: parsed = parse_one(sql, dialect=dialect) # 遍历所有JOIN for join in parsed.find_all(exp.Join): # 检查是否有ON条件 on_condition = join.args.get("on") if on_condition is None: # 检查是否使用USING using = join.args.get("using") if using is None: issues.append("JOIN缺少ON条件") else: # ON条件为空表达式 if isinstance(on_condition, exp.Empty): issues.append("JOIN的ON条件为空") except Exception as e: logger.debug(f"JOIN条件检查异常: {e}") return issues def validate_order_by_fields( sql: str, schema_manager, dialect: str = "tsql" ) -> List[str]: """ 验证ORDER BY字段是否存在于对应表中 Args: sql: SQL语句 schema_manager: Schema管理器 dialect: SQL方言 Returns: 问题列表 """ from utils.sql_parser import parse_one, exp issues = [] try: parsed = parse_one(sql, dialect=dialect) order = parsed.args.get("order") if order: for ordered in order.expressions: expr = ordered.this # 提取字段和表 if isinstance(expr, exp.Column): tbl_name = expr.table col_name = expr.name if tbl_name: table = schema_manager.get_table(tbl_name) if table: col_names = [c.name for c in table.columns] if col_name not in col_names: issues.append( f"ORDER BY字段不存在: {tbl_name}.{col_name}" ) except Exception as e: logger.debug(f"ORDER BY验证异常: {e}") return issues def full_validation_pipeline( sql: str, schema_manager, dialect: str = "tsql", check_dangerous: bool = True ) -> Dict: """ 完整验证流水线 Args: sql: SQL语句 schema_manager: Schema管理器 dialect: SQL方言 check_dangerous: 是否检查危险操作 Returns: 验证结果字典 """ result = { "valid": True, "errors": [], "warnings": [], "suggestions": [] } # 1. 语法验证 syntax_ok, syntax_errors = validate_sql_syntax(sql, dialect) if not syntax_ok: result["valid"] = False result["errors"].extend(syntax_errors) # 2. 危险操作检查 if check_dangerous: safe, dangers = check_dangerous_operations(sql) if not safe: result["valid"] = False result["errors"].append(f"包含危险操作: {', '.join(dangers)}") # 3. Schema一致性验证 schema_ok, schema_errors = validate_schema_consistency(sql, schema_manager, dialect) if not schema_ok: result["valid"] = False result["errors"].extend(schema_errors) # 4. GROUP BY验证 groupby_ok, groupby_errors = validate_aggregation_groupby(sql, dialect) if not groupby_ok: result["valid"] = False result["errors"].extend(groupby_errors) # 5. JOIN条件验证 join_issues = check_join_conditions(sql, dialect) if join_issues: result["valid"] = False result["errors"].extend(join_issues) # 6. ORDER BY验证 order_issues = validate_order_by_fields(sql, schema_manager, dialect) if order_issues: result["warnings"].extend(order_issues) # 7. SQL注入模式检查(警告) injection_patterns = check_sql_injection_patterns(sql) if injection_patterns: result["warnings"].append(f"检测到可疑模式: {', '.join(injection_patterns)}") return result