347 lines
9.9 KiB
Python
347 lines
9.9 KiB
Python
"""
|
||||
|
|
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
|