Files
ai-g3sb-backman2.0/backend/utils/validators.py
T

509 lines
16 KiB
Python
Raw Normal View History

2026-04-10 16:52:07 +08:00
"""
SQL 验证工具集
"""
import re
import logging
from typing import Tuple, List, Dict, Optional
2026-04-10 16:52:07 +08:00
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 _cjk_text_in_string_literal(text: str) -> bool:
return bool(_CJK_IN_STRING_RE.search(text))
2026-04-14 10:28:22 +08:00
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 即失败。
跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本。
2026-04-14 10:28:22 +08:00
"""
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):
2026-04-14 10:28:22 +08:00
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):
2026-04-14 10:28:22 +08:00
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
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