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
|