2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
SQL 验证工具集
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
import re
|
|
|
|
|
|
import logging
|
|
|
|
|
|
from typing import Tuple, List, Dict
|
|
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# CJK Unified Ideographs + 兼容扩展(用于禁止中文业务词出现在 SQL 字符串字面量中)
|
|
|
|
|
|
_CJK_IN_STRING_RE = re.compile(
|
|
|
|
|
|
r"[\u3000-\u303f\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff\uff00-\uffef]"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
扫描 SQL 中单引号字符串(含 T-SQL 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]:
|
|
|
|
|
|
"""从 start 指向的 opening `'` 之后开始读,返回 (内容, 闭合引号后下标)。"""
|
|
|
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
|
# N' 或 n' 前缀的 Unicode 字面量
|
|
|
|
|
|
if i + 1 < n and sql[i] in "Nn" and sql[i + 1] == "'":
|
|
|
|
|
|
body, i = _read_single_quoted_string(i + 2)
|
|
|
|
|
|
if _CJK_IN_STRING_RE.search(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_IN_STRING_RE.search(body):
|
|
|
|
|
|
prev = body[:48] + ("…" if len(body) > 48 else "")
|
|
|
|
|
|
errors.append(
|
|
|
|
|
|
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
|
|
|
|
|
|
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
|
|
|
|
|
|
)
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
i += 1
|
|
|
|
|
|
|
|
|
|
|
|
return len(errors) == 0, errors
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
# 危险操作关键词(除非明确允许)
|
|
|
|
|
|
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
|