""" SQL 验证工具集 """ import re import logging from typing import Tuple, List, Dict, Optional logger = logging.getLogger(__name__) # CJK Unified Ideographs + 兼容扩展(用于禁止中文业务词出现在 SQL 字符串字面量中) _CJK_IN_STRING_RE = re.compile( r"[\u3000-\u303f\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff\uff00-\uffef]" ) def _cjk_text_in_string_literal(text: str) -> bool: return bool(_CJK_IN_STRING_RE.search(text)) def _node_in_subtree(root, target) -> bool: if root is None: return False for n in root.walk(): if n is target: return True return False def _cjk_string_in_forbidden_context(node, exp) -> bool: """ 禁止含 CJK 的字面量出现在「比对/过滤」语境:WHERE、HAVING、JOIN ON、 以及 CASE 分支的 WHEN 条件(含简单 CASE 的 WHEN 值), 但允许出现在 SELECT 投影、CASE 的 THEN/ELSE 结果等纯展示位置。 """ if node.find_ancestor(exp.Where): return True if node.find_ancestor(exp.Having): return True join = node.find_ancestor(exp.Join) if join is not None: on = join.args.get("on") if on is not None and _node_in_subtree(on, node): return True case = node.find_ancestor(exp.Case) while case is not None: default = case.args.get("default") if default is not None and _node_in_subtree(default, node): return False for br in case.args.get("ifs") or []: then_expr = br.args.get("true") cond = br.this if then_expr is not None and _node_in_subtree(then_expr, node): return False if cond is not None and _node_in_subtree(cond, node): return True case = case.find_ancestor(exp.Case) return False def _check_no_cjk_legacy_text_scan(sql: str) -> Tuple[bool, List[str]]: """ 解析失败时的回退:扫描单引号字符串(含 N'…'),字面量内出现 CJK 即失败。 跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本。 """ errors: List[str] = [] i = 0 n = len(sql) in_line_comment = False in_block_comment = False def _read_single_quoted_string(start: int) -> Tuple[str, int]: j = start parts: List[str] = [] while j < n: ch = sql[j] if ch == "'": if j + 1 < n and sql[j + 1] == "'": parts.append("'") j += 2 continue return "".join(parts), j + 1 parts.append(ch) j += 1 return "".join(parts), j while i < n: if in_line_comment: if sql[i] == "\n": in_line_comment = False i += 1 continue if in_block_comment: if i + 1 < n and sql[i : i + 2] == "*/": in_block_comment = False i += 2 else: i += 1 continue two = sql[i : i + 2] if two == "--": in_line_comment = True i += 2 continue if two == "/*": in_block_comment = True i += 2 continue if i + 1 < n and sql[i] in "Nn" and sql[i + 1] == "'": body, i = _read_single_quoted_string(i + 2) if _cjk_text_in_string_literal(body): prev = body[:48] + ("…" if len(body) > 48 else "") errors.append( "SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。" f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。" ) continue if sql[i] == "'": body, i = _read_single_quoted_string(i + 1) if _cjk_text_in_string_literal(body): prev = body[:48] + ("…" if len(body) > 48 else "") errors.append( "SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。" f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。" ) continue i += 1 return len(errors) == 0, errors def check_no_cjk_in_sql_string_literals(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]: """ 禁止在「过滤/比对」语境使用含中日韩字符的字符串字面量(含 T-SQL ``N'…'``)。 允许在 SELECT 投影、CASE 的 THEN/ELSE 结果等展示用字面量中使用中文标签。 解析失败时回退为全文扫描(与旧版一致,偏严)。 """ from sqlglot import parse_one, exp errors: List[str] = [] try: parsed = parse_one(sql, dialect=dialect) except Exception as e: logger.debug("CJK 校验回退为全文扫描(SQL 解析失败): %s", e) return _check_no_cjk_legacy_text_scan(sql) for node in parsed.walk(): text: Optional[str] = None if isinstance(node, exp.Literal) and node.is_string: text = str(node.this) elif isinstance(node, exp.National): text = str(node.this) if text is None or not _cjk_text_in_string_literal(text): continue if _cjk_string_in_forbidden_context(node, exp): prev = text[:48] + ("…" if len(text) > 48 else "") errors.append( "SQL 在 WHERE/HAVING/JOIN ON 或 CASE/WHEN 条件中出现含中文的字符串字面量(禁止)。" f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。" ) return len(errors) == 0, errors # 危险操作关键词(除非明确允许) 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