first commit
This commit is contained in:
@@ -0,0 +1,346 @@
|
||||
"""
|
||||
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
|
||||
Reference in New Issue
Block a user