springbootAI 1.8.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- spring/__init__.py +66 -0
- spring/ai/__init__.py +78 -0
- spring/ai/advisors.py +139 -0
- spring/ai/annotations.py +74 -0
- spring/ai/autoconfig.py +481 -0
- spring/ai/core.py +391 -0
- spring/ai/etl.py +188 -0
- spring/ai/memory.py +109 -0
- spring/ai/observability.py +129 -0
- spring/ai/providers.py +789 -0
- spring/ai/resilience.py +258 -0
- spring/ai/tools.py +106 -0
- spring/ai/vectorstore.py +303 -0
- spring/annotations/__init__.py +188 -0
- spring/annotations/cache.py +126 -0
- spring/annotations/cloud.py +207 -0
- spring/annotations/conditional.py +272 -0
- spring/annotations/core.py +864 -0
- spring/annotations/messaging.py +107 -0
- spring/aop/__init__.py +4 -0
- spring/aop/cloud_aop.py +404 -0
- spring/aop/comprehensive_aop.py +1015 -0
- spring/aop/method_interceptor.py +19 -0
- spring/aop/proxy_factory.py +55 -0
- spring/cloud/__init__.py +76 -0
- spring/cloud/discovery.py +364 -0
- spring/cloud/feign.py +469 -0
- spring/cloud/gateway.py +452 -0
- spring/cloud/load_balancer.py +149 -0
- spring/cloud/seata.py +557 -0
- spring/cloud/sentinel.py +525 -0
- spring/cloud/tracer.py +337 -0
- spring/config/__init__.py +21 -0
- spring/config/binding.py +206 -0
- spring/config/config_loader.py +405 -0
- spring/context/__init__.py +13 -0
- spring/context/application_context.py +589 -0
- spring/context/bean_definition.py +70 -0
- spring/context/bean_factory.py +1052 -0
- spring/context/registry.py +58 -0
- spring/context/scanner.py +106 -0
- spring/core/__init__.py +3 -0
- spring/core/graceful_shutdown.py +196 -0
- spring/core/typing_utils.py +50 -0
- spring/csv/__init__.py +52 -0
- spring/csv/annotations.py +402 -0
- spring/csv/converters.py +69 -0
- spring/csv/easy_csv.py +95 -0
- spring/csv/exceptions.py +27 -0
- spring/csv/reader.py +195 -0
- spring/csv/writer.py +155 -0
- spring/data/__init__.py +54 -0
- spring/data/page.py +181 -0
- spring/data/repository.py +274 -0
- spring/data/specification.py +228 -0
- spring/datasource/__init__.py +66 -0
- spring/datasource/annotations.py +133 -0
- spring/datasource/context.py +69 -0
- spring/datasource/dynamic.py +148 -0
- spring/event/__init__.py +7 -0
- spring/event/publisher.py +69 -0
- spring/excel/__init__.py +51 -0
- spring/excel/annotations.py +405 -0
- spring/excel/converters.py +231 -0
- spring/excel/easy_excel.py +94 -0
- spring/excel/exceptions.py +31 -0
- spring/excel/reader.py +254 -0
- spring/excel/style.py +95 -0
- spring/excel/writer.py +197 -0
- spring/i18n/__init__.py +97 -0
- spring/i18n/accessor.py +94 -0
- spring/i18n/auto_config.py +177 -0
- spring/i18n/holder.py +106 -0
- spring/i18n/locale.py +152 -0
- spring/i18n/locale_resolver.py +367 -0
- spring/i18n/message_source.py +250 -0
- spring/i18n/middleware.py +79 -0
- spring/i18n/properties.py +168 -0
- spring/i18n/sources.py +255 -0
- spring/logging/__init__.py +1 -0
- spring/logging/loguru_logger.py +228 -0
- spring/main.py +378 -0
- spring/messaging/__init__.py +1 -0
- spring/messaging/rabbitmq.py +302 -0
- spring/monitoring/__init__.py +1 -0
- spring/monitoring/prometheus.py +199 -0
- spring/orm/__init__.py +258 -0
- spring/orm/database.py +222 -0
- spring/orm/ddl_auto.py +1217 -0
- spring/orm/migration.py +419 -0
- spring/orm/mybatis_integration.py +400 -0
- spring/orm/pymybatis/__init__.py +86 -0
- spring/orm/pymybatis/annotations/__init__.py +30 -0
- spring/orm/pymybatis/annotations/annotations.py +332 -0
- spring/orm/pymybatis/cache/__init__.py +47 -0
- spring/orm/pymybatis/cache/cache.py +371 -0
- spring/orm/pymybatis/cache/redis_cache.py +434 -0
- spring/orm/pymybatis/circuit_breaker/__init__.py +21 -0
- spring/orm/pymybatis/circuit_breaker/circuit_breaker.py +424 -0
- spring/orm/pymybatis/configuration.py +525 -0
- spring/orm/pymybatis/core/__init__.py +10 -0
- spring/orm/pymybatis/core/sql_session.py +1382 -0
- spring/orm/pymybatis/core/sql_session_factory.py +76 -0
- spring/orm/pymybatis/dialect/__init__.py +9 -0
- spring/orm/pymybatis/dialect/dialect.py +445 -0
- spring/orm/pymybatis/dynamic_sql/__init__.py +9 -0
- spring/orm/pymybatis/dynamic_sql/dynamic_sql.py +900 -0
- spring/orm/pymybatis/interceptor/__init__.py +31 -0
- spring/orm/pymybatis/interceptor/interceptor.py +427 -0
- spring/orm/pymybatis/mapper/__init__.py +9 -0
- spring/orm/pymybatis/mapper/mapper.py +540 -0
- spring/orm/pymybatis/metrics/__init__.py +41 -0
- spring/orm/pymybatis/metrics/metrics.py +595 -0
- spring/orm/pymybatis/pool/__init__.py +9 -0
- spring/orm/pymybatis/pool/connection_pool.py +711 -0
- spring/orm/pymybatis/security/__init__.py +19 -0
- spring/orm/pymybatis/security/access_control.py +415 -0
- spring/orm/pymybatis/security/password_encoder.py +293 -0
- spring/orm/pymybatis/security/sensitive_data_masker.py +326 -0
- spring/orm/pymybatis/security/sql_injection_detector.py +675 -0
- spring/orm/pymybatis/transaction/__init__.py +9 -0
- spring/orm/pymybatis/transaction/transaction.py +288 -0
- spring/orm/pymybatis/type_handler/__init__.py +37 -0
- spring/orm/pymybatis/type_handler/type_handler.py +473 -0
- spring/orm/pymybatis/version.py +9 -0
- spring/orm/pymybatis/xml_parser/__init__.py +9 -0
- spring/orm/pymybatis/xml_parser/xml_parser.py +761 -0
- spring/retry/__init__.py +12 -0
- spring/retry/retry_annotations.py +71 -0
- spring/retry/retry_decorator.py +155 -0
- spring/scheduling/__init__.py +3 -0
- spring/scheduling/scheduler.py +389 -0
- spring/security/__init__.py +39 -0
- spring/security/jwt_utils.py +281 -0
- spring/security/replay_protection.py +206 -0
- spring/security/secret_manager.py +226 -0
- spring/security/security_aop.py +248 -0
- spring/security/security_context.py +172 -0
- spring/test/__init__.py +45 -0
- spring/test/slicing.py +341 -0
- spring/tracing/__init__.py +11 -0
- spring/tracing/skywalking.py +229 -0
- spring/tx/__init__.py +52 -0
- spring/tx/events.py +172 -0
- spring/tx/synchronization.py +143 -0
- spring/utils/__init__.py +5 -0
- spring/utils/banner.py +32 -0
- spring/utils/logger.py +73 -0
- spring/utils/redis_client.py +526 -0
- spring/validation/__init__.py +55 -0
- spring/validation/aop.py +141 -0
- spring/validation/constraints.py +357 -0
- spring/validation/exceptions.py +55 -0
- spring/validation/validator.py +139 -0
- spring/web/__init__.py +12 -0
- spring/web/actuator.py +319 -0
- spring/web/exception_handler.py +61 -0
- spring/web/health.py +399 -0
- spring/web/interceptor.py +91 -0
- spring/web/result.py +44 -0
- spring/web/swagger.py +601 -0
- spring/web/web_context.py +755 -0
- spring/websocket/__init__.py +86 -0
- spring/websocket/annotations.py +169 -0
- spring/websocket/broker.py +238 -0
- spring/websocket/exceptions.py +26 -0
- spring/websocket/handler.py +243 -0
- spring/websocket/router.py +526 -0
- spring/websocket/session.py +216 -0
- springbootai-1.8.0.dist-info/METADATA +2796 -0
- springbootai-1.8.0.dist-info/RECORD +175 -0
- springbootai-1.8.0.dist-info/WHEEL +5 -0
- springbootai-1.8.0.dist-info/entry_points.txt +2 -0
- springbootai-1.8.0.dist-info/licenses/LICENSE +7 -0
- springbootai-1.8.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,675 @@
|
|
|
1
|
+
"""
|
|
2
|
+
PyMyBatis SQL注入检测模块
|
|
3
|
+
|
|
4
|
+
实现SQL注入攻击的主动检测和防御机制,核心安全特性:
|
|
5
|
+
- 参数值注入检测(正则 + AST双重验证)
|
|
6
|
+
- SQL语句注入检测
|
|
7
|
+
- DDL语句禁用(DROP/ALTER/CREATE/TRUNCATE等)
|
|
8
|
+
- ${}参数白名单检查(支持表名/字段名白名单)
|
|
9
|
+
- AST解析验证(基于sqlglot,可选)
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
import re
|
|
13
|
+
import logging
|
|
14
|
+
from typing import Optional, Any, Dict, Set, List
|
|
15
|
+
from enum import Enum
|
|
16
|
+
|
|
17
|
+
logger = logging.getLogger(__name__)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class SQLInjectionLevel(Enum):
|
|
21
|
+
"""SQL注入风险级别"""
|
|
22
|
+
NONE = 0
|
|
23
|
+
LOW = 1
|
|
24
|
+
MEDIUM = 2
|
|
25
|
+
HIGH = 3
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class SQLInjectionPattern:
|
|
29
|
+
"""SQL注入模式定义"""
|
|
30
|
+
|
|
31
|
+
def __init__(self, pattern: str, level: SQLInjectionLevel, description: str):
|
|
32
|
+
self.pattern = re.compile(pattern, re.IGNORECASE)
|
|
33
|
+
self.level = level
|
|
34
|
+
self.description = description
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class SQLInjectionDetector:
|
|
38
|
+
"""
|
|
39
|
+
SQL注入检测器
|
|
40
|
+
|
|
41
|
+
核心功能:
|
|
42
|
+
1. 参数值注入检测(正则 + AST双重验证)
|
|
43
|
+
2. SQL语句注入检测
|
|
44
|
+
3. DDL语句检测与禁用
|
|
45
|
+
4. ${}参数安全检查
|
|
46
|
+
5. 返回风险等级和建议
|
|
47
|
+
"""
|
|
48
|
+
|
|
49
|
+
# SQL注入模式列表
|
|
50
|
+
INJECTION_PATTERNS = [
|
|
51
|
+
# 基于注释的注入
|
|
52
|
+
SQLInjectionPattern(
|
|
53
|
+
r'(--(?:\s|$)|(?:^|\s)#|/\*).*',
|
|
54
|
+
SQLInjectionLevel.HIGH,
|
|
55
|
+
'注释注入',
|
|
56
|
+
),
|
|
57
|
+
|
|
58
|
+
# UNION注入
|
|
59
|
+
SQLInjectionPattern(r'\bUNION\b.*\bSELECT\b', SQLInjectionLevel.HIGH, 'UNION注入'),
|
|
60
|
+
|
|
61
|
+
# 布尔盲注
|
|
62
|
+
SQLInjectionPattern(r'\b(AND|OR)\b.*\d+\s*(=|<|>)\s*\d+', SQLInjectionLevel.HIGH, '布尔盲注'),
|
|
63
|
+
|
|
64
|
+
# 简单字符串注入 (如 ' OR '1'='1)
|
|
65
|
+
SQLInjectionPattern(r"'.*\s*(OR|AND)\s*'.*'.*='.*'", SQLInjectionLevel.HIGH, '字符串注入'),
|
|
66
|
+
SQLInjectionPattern(r"'.*\s*(OR|AND)\s*\d+\s*=\s*\d+", SQLInjectionLevel.HIGH, '数字条件注入'),
|
|
67
|
+
SQLInjectionPattern(r"'.*\s*(OR|AND)\s*1\s*=\s*1", SQLInjectionLevel.HIGH, '恒真条件注入'),
|
|
68
|
+
SQLInjectionPattern(r"1'\s*OR\s*'1'\s*=\s*'1", SQLInjectionLevel.HIGH, '经典字符串注入'),
|
|
69
|
+
SQLInjectionPattern(r"'.*\s*(OR|AND)\s*\d+\s*=\s*\d+.*'", SQLInjectionLevel.HIGH, '闭合注入'),
|
|
70
|
+
|
|
71
|
+
# 时间盲注
|
|
72
|
+
SQLInjectionPattern(r'\b(SLEEP|BENCHMARK|WAITFOR)\b', SQLInjectionLevel.HIGH, '时间盲注'),
|
|
73
|
+
|
|
74
|
+
# 基于函数的注入
|
|
75
|
+
SQLInjectionPattern(
|
|
76
|
+
r'\b(CONCAT|GROUP_CONCAT|VERSION|DATABASE|USER)\s*\(',
|
|
77
|
+
SQLInjectionLevel.MEDIUM,
|
|
78
|
+
'信息收集函数',
|
|
79
|
+
),
|
|
80
|
+
|
|
81
|
+
# 危险关键字
|
|
82
|
+
SQLInjectionPattern(r'\b(DROP|DELETE|UPDATE|INSERT|TRUNCATE|ALTER|CREATE|GRANT|REVOKE)\b',
|
|
83
|
+
SQLInjectionLevel.HIGH, '危险SQL关键字'),
|
|
84
|
+
|
|
85
|
+
# 基于子查询的注入
|
|
86
|
+
SQLInjectionPattern(r'\(\s*SELECT\s+', SQLInjectionLevel.MEDIUM, '子查询注入'),
|
|
87
|
+
|
|
88
|
+
# 基于编码的注入
|
|
89
|
+
SQLInjectionPattern(r'(0x[0-9a-f]+|char\(|ascii\()', SQLInjectionLevel.MEDIUM, '编码注入'),
|
|
90
|
+
|
|
91
|
+
# 换行符注入
|
|
92
|
+
SQLInjectionPattern(r'[\r\n]+\s*(SELECT|INSERT|UPDATE|DELETE|DROP)', SQLInjectionLevel.HIGH, '换行注入'),
|
|
93
|
+
|
|
94
|
+
# 多个连续空格(可能是绕过尝试)
|
|
95
|
+
SQLInjectionPattern(r'\s{3,}', SQLInjectionLevel.LOW, '异常空格'),
|
|
96
|
+
|
|
97
|
+
# 字符串拼接
|
|
98
|
+
SQLInjectionPattern(r"('.*')\s*(\|\||\+)\s*('.*')", SQLInjectionLevel.MEDIUM, '字符串拼接'),
|
|
99
|
+
|
|
100
|
+
# 条件注释
|
|
101
|
+
SQLInjectionPattern(r'/\*\s*!\d+\s*', SQLInjectionLevel.HIGH, 'MySQL条件注释'),
|
|
102
|
+
|
|
103
|
+
# 执行命令
|
|
104
|
+
SQLInjectionPattern(r'\b(EXEC|EXECUTE|XP_CMDSHELL|SYSTEM|SHELL)\b', SQLInjectionLevel.HIGH, '命令执行'),
|
|
105
|
+
|
|
106
|
+
# 堆叠查询
|
|
107
|
+
SQLInjectionPattern(r';\s*(SELECT|INSERT|UPDATE|DELETE|DROP)', SQLInjectionLevel.HIGH, '堆叠查询'),
|
|
108
|
+
|
|
109
|
+
# 回显注入
|
|
110
|
+
SQLInjectionPattern(r'\b(CAST|CONVERT)\b.*\b(VARCHAR|CHAR)\b', SQLInjectionLevel.MEDIUM, '类型转换'),
|
|
111
|
+
|
|
112
|
+
# 正则注入
|
|
113
|
+
SQLInjectionPattern(r'\bREGEXP\b.*\'.*\'', SQLInjectionLevel.MEDIUM, '正则注入'),
|
|
114
|
+
]
|
|
115
|
+
|
|
116
|
+
# DDL语句模式(生产环境默认禁用)
|
|
117
|
+
DDL_PATTERNS = [
|
|
118
|
+
SQLInjectionPattern(r'^\s*DROP\s+', SQLInjectionLevel.HIGH, 'DROP语句'),
|
|
119
|
+
SQLInjectionPattern(r'^\s*ALTER\s+', SQLInjectionLevel.HIGH, 'ALTER语句'),
|
|
120
|
+
SQLInjectionPattern(r'^\s*CREATE\s+(TABLE|INDEX|VIEW|FUNCTION|PROCEDURE)', SQLInjectionLevel.HIGH, 'CREATE语句'),
|
|
121
|
+
SQLInjectionPattern(r'^\s*TRUNCATE\s+', SQLInjectionLevel.HIGH, 'TRUNCATE语句'),
|
|
122
|
+
SQLInjectionPattern(r'^\s*RENAME\s+', SQLInjectionLevel.HIGH, 'RENAME语句'),
|
|
123
|
+
SQLInjectionPattern(r'^\s*GRANT\s+', SQLInjectionLevel.HIGH, 'GRANT语句'),
|
|
124
|
+
SQLInjectionPattern(r'^\s*REVOKE\s+', SQLInjectionLevel.HIGH, 'REVOKE语句'),
|
|
125
|
+
SQLInjectionPattern(r'^\s*COMMIT\s+', SQLInjectionLevel.MEDIUM, 'COMMIT语句'),
|
|
126
|
+
SQLInjectionPattern(r'^\s*ROLLBACK\s+', SQLInjectionLevel.MEDIUM, 'ROLLBACK语句'),
|
|
127
|
+
]
|
|
128
|
+
|
|
129
|
+
# ${}参数白名单(表名、字段名等)
|
|
130
|
+
RAW_PARAM_WHITELIST_PATTERNS = [
|
|
131
|
+
re.compile(r'^[a-zA-Z_][a-zA-Z0-9_]*$'), # 表名/字段名
|
|
132
|
+
re.compile(r'^(ASC|DESC)$'), # 排序方向
|
|
133
|
+
re.compile(r'^[a-zA-Z_][a-zA-Z0-9_.]*$'), # 带schema的表名
|
|
134
|
+
]
|
|
135
|
+
|
|
136
|
+
def __init__(self, enabled: bool = True,
|
|
137
|
+
max_risk_level: SQLInjectionLevel = SQLInjectionLevel.LOW,
|
|
138
|
+
block_ddl: bool = True,
|
|
139
|
+
allow_raw_params: bool = False,
|
|
140
|
+
raw_param_whitelist: Optional[Set[str]] = None,
|
|
141
|
+
allowed_tables: Optional[Set[str]] = None,
|
|
142
|
+
allowed_columns: Optional[Set[str]] = None,
|
|
143
|
+
enable_ast_validation: bool = False):
|
|
144
|
+
"""
|
|
145
|
+
初始化SQL注入检测器
|
|
146
|
+
|
|
147
|
+
Args:
|
|
148
|
+
enabled: 是否启用检测
|
|
149
|
+
max_risk_level: 允许的最大风险级别,超过此级别会阻止执行
|
|
150
|
+
block_ddl: 是否阻止DDL语句
|
|
151
|
+
allow_raw_params: 是否允许${}参数
|
|
152
|
+
raw_param_whitelist: ${}参数名白名单
|
|
153
|
+
allowed_tables: 允许的表名白名单
|
|
154
|
+
allowed_columns: 允许的字段名白名单
|
|
155
|
+
enable_ast_validation: 是否启用AST验证(需要sqlglot库)
|
|
156
|
+
"""
|
|
157
|
+
self.enabled = enabled
|
|
158
|
+
self.max_risk_level = max_risk_level
|
|
159
|
+
self.block_ddl = block_ddl
|
|
160
|
+
self.allow_raw_params = allow_raw_params
|
|
161
|
+
self.raw_param_whitelist = raw_param_whitelist or set()
|
|
162
|
+
self.allowed_tables = allowed_tables or set()
|
|
163
|
+
self.allowed_columns = allowed_columns or set()
|
|
164
|
+
self.enable_ast_validation = enable_ast_validation
|
|
165
|
+
|
|
166
|
+
# 延迟加载sqlglot
|
|
167
|
+
self._sqlglot = None
|
|
168
|
+
|
|
169
|
+
def _load_sqlglot(self):
|
|
170
|
+
"""延迟加载sqlglot库"""
|
|
171
|
+
if self._sqlglot is None:
|
|
172
|
+
try:
|
|
173
|
+
import sqlglot
|
|
174
|
+
self._sqlglot = sqlglot
|
|
175
|
+
except ImportError:
|
|
176
|
+
logger.warning("sqlglot库未安装,AST验证功能将不可用。安装方法: pip install sqlglot")
|
|
177
|
+
self.enable_ast_validation = False
|
|
178
|
+
return self._sqlglot
|
|
179
|
+
|
|
180
|
+
def detect(self, value: Any) -> SQLInjectionLevel:
|
|
181
|
+
"""
|
|
182
|
+
检测单个值是否包含SQL注入风险
|
|
183
|
+
|
|
184
|
+
Args:
|
|
185
|
+
value: 待检测的值
|
|
186
|
+
|
|
187
|
+
Returns:
|
|
188
|
+
风险等级
|
|
189
|
+
"""
|
|
190
|
+
if not self.enabled:
|
|
191
|
+
return SQLInjectionLevel.NONE
|
|
192
|
+
|
|
193
|
+
if value is None:
|
|
194
|
+
return SQLInjectionLevel.NONE
|
|
195
|
+
|
|
196
|
+
# 转换为字符串进行检测
|
|
197
|
+
if not isinstance(value, str):
|
|
198
|
+
value = str(value)
|
|
199
|
+
|
|
200
|
+
# 首先使用正则检测
|
|
201
|
+
max_level = self._detect_by_regex(value)
|
|
202
|
+
|
|
203
|
+
# 如果正则检测到风险或启用了AST验证,进行AST验证
|
|
204
|
+
if max_level != SQLInjectionLevel.NONE or self.enable_ast_validation:
|
|
205
|
+
ast_level = self._detect_by_ast(value)
|
|
206
|
+
if ast_level.value > max_level.value:
|
|
207
|
+
max_level = ast_level
|
|
208
|
+
|
|
209
|
+
return max_level
|
|
210
|
+
|
|
211
|
+
def _detect_by_regex(self, value: str) -> SQLInjectionLevel:
|
|
212
|
+
"""
|
|
213
|
+
使用正则表达式检测SQL注入风险
|
|
214
|
+
|
|
215
|
+
Args:
|
|
216
|
+
value: 待检测的字符串
|
|
217
|
+
|
|
218
|
+
Returns:
|
|
219
|
+
风险等级
|
|
220
|
+
"""
|
|
221
|
+
max_level = SQLInjectionLevel.NONE
|
|
222
|
+
|
|
223
|
+
for pattern in self.INJECTION_PATTERNS:
|
|
224
|
+
if pattern.pattern.search(value):
|
|
225
|
+
if pattern.level.value > max_level.value:
|
|
226
|
+
max_level = pattern.level
|
|
227
|
+
# 一旦检测到最高级别,立即返回
|
|
228
|
+
if max_level == SQLInjectionLevel.HIGH:
|
|
229
|
+
return max_level
|
|
230
|
+
|
|
231
|
+
return max_level
|
|
232
|
+
|
|
233
|
+
def _detect_by_ast(self, value: str) -> SQLInjectionLevel:
|
|
234
|
+
"""
|
|
235
|
+
使用AST解析检测SQL注入风险
|
|
236
|
+
|
|
237
|
+
Args:
|
|
238
|
+
value: 待检测的字符串
|
|
239
|
+
|
|
240
|
+
Returns:
|
|
241
|
+
风险等级
|
|
242
|
+
"""
|
|
243
|
+
if not self.enable_ast_validation:
|
|
244
|
+
return SQLInjectionLevel.NONE
|
|
245
|
+
|
|
246
|
+
sqlglot = self._load_sqlglot()
|
|
247
|
+
if sqlglot is None:
|
|
248
|
+
return SQLInjectionLevel.NONE
|
|
249
|
+
|
|
250
|
+
try:
|
|
251
|
+
# 尝试解析为SQL表达式
|
|
252
|
+
parsed = sqlglot.parse_one(value)
|
|
253
|
+
|
|
254
|
+
# 检查AST结构中的危险模式
|
|
255
|
+
return self._analyze_ast(parsed)
|
|
256
|
+
except Exception as e:
|
|
257
|
+
# 解析失败,可能是恶意输入
|
|
258
|
+
logger.debug(f"AST解析失败,可能存在注入风险: {e}")
|
|
259
|
+
return SQLInjectionLevel.MEDIUM
|
|
260
|
+
|
|
261
|
+
def _analyze_ast(self, parsed) -> SQLInjectionLevel:
|
|
262
|
+
"""
|
|
263
|
+
分析AST结构,检测危险模式
|
|
264
|
+
|
|
265
|
+
Args:
|
|
266
|
+
parsed: sqlglot解析后的AST节点
|
|
267
|
+
|
|
268
|
+
Returns:
|
|
269
|
+
风险等级
|
|
270
|
+
"""
|
|
271
|
+
sqlglot = self._sqlglot
|
|
272
|
+
if sqlglot is None:
|
|
273
|
+
return SQLInjectionLevel.NONE
|
|
274
|
+
|
|
275
|
+
max_level = SQLInjectionLevel.NONE
|
|
276
|
+
|
|
277
|
+
# 遍历AST节点
|
|
278
|
+
for node in parsed.walk():
|
|
279
|
+
node_type = type(node).__name__
|
|
280
|
+
|
|
281
|
+
# 检测UNION注入
|
|
282
|
+
if node_type == 'Union':
|
|
283
|
+
return SQLInjectionLevel.HIGH
|
|
284
|
+
|
|
285
|
+
# 检测子查询
|
|
286
|
+
if node_type == 'Subquery':
|
|
287
|
+
max_level = SQLInjectionLevel.MEDIUM
|
|
288
|
+
|
|
289
|
+
# 检测危险函数
|
|
290
|
+
if node_type == 'Func':
|
|
291
|
+
func_name = node.this.this.lower() if hasattr(node.this, 'this') else ''
|
|
292
|
+
dangerous_functions = ['sleep', 'benchmark', 'waitfor', 'version', 'database', 'user', 'system']
|
|
293
|
+
if func_name in dangerous_functions:
|
|
294
|
+
return SQLInjectionLevel.HIGH
|
|
295
|
+
|
|
296
|
+
# 检测DDL语句
|
|
297
|
+
if node_type in ['Drop', 'Truncate', 'Alter', 'Create']:
|
|
298
|
+
return SQLInjectionLevel.HIGH
|
|
299
|
+
|
|
300
|
+
return max_level
|
|
301
|
+
|
|
302
|
+
def detect_ddl(self, sql: str) -> bool:
|
|
303
|
+
"""
|
|
304
|
+
检测SQL语句是否为DDL语句(结合正则和AST)
|
|
305
|
+
|
|
306
|
+
Args:
|
|
307
|
+
sql: SQL语句
|
|
308
|
+
|
|
309
|
+
Returns:
|
|
310
|
+
是否为DDL语句
|
|
311
|
+
"""
|
|
312
|
+
if not self.block_ddl:
|
|
313
|
+
return False
|
|
314
|
+
|
|
315
|
+
if sql is None:
|
|
316
|
+
return False
|
|
317
|
+
|
|
318
|
+
if not isinstance(sql, str):
|
|
319
|
+
sql = str(sql)
|
|
320
|
+
|
|
321
|
+
# 首先使用正则检测
|
|
322
|
+
for pattern in self.DDL_PATTERNS:
|
|
323
|
+
if pattern.pattern.search(sql.strip()):
|
|
324
|
+
logger.warning(f"检测到DDL语句: {sql}")
|
|
325
|
+
return True
|
|
326
|
+
|
|
327
|
+
# 如果启用了AST验证,进行二次验证
|
|
328
|
+
if self.enable_ast_validation:
|
|
329
|
+
sqlglot = self._load_sqlglot()
|
|
330
|
+
if sqlglot is not None:
|
|
331
|
+
try:
|
|
332
|
+
parsed = sqlglot.parse_one(sql)
|
|
333
|
+
for node in parsed.walk():
|
|
334
|
+
node_type = type(node).__name__
|
|
335
|
+
if node_type in ['Drop', 'Truncate', 'Alter', 'Create', 'Grant', 'Revoke']:
|
|
336
|
+
logger.warning(f"AST检测到DDL语句: {sql}")
|
|
337
|
+
return True
|
|
338
|
+
except Exception:
|
|
339
|
+
pass
|
|
340
|
+
|
|
341
|
+
return False
|
|
342
|
+
|
|
343
|
+
def is_ddl_blocked(self, sql: str) -> bool:
|
|
344
|
+
"""
|
|
345
|
+
判断DDL语句是否被阻止
|
|
346
|
+
|
|
347
|
+
Args:
|
|
348
|
+
sql: SQL语句
|
|
349
|
+
|
|
350
|
+
Returns:
|
|
351
|
+
是否被阻止
|
|
352
|
+
"""
|
|
353
|
+
return self.block_ddl and self.detect_ddl(sql)
|
|
354
|
+
|
|
355
|
+
def detect_raw_param(self, param_name: str, param_value: Any, param_type: Optional[str] = None) -> bool:
|
|
356
|
+
"""
|
|
357
|
+
检测${}参数是否安全(增强版:支持表名/字段名白名单)
|
|
358
|
+
|
|
359
|
+
Args:
|
|
360
|
+
param_name: 参数名
|
|
361
|
+
param_value: 参数值
|
|
362
|
+
param_type: 参数类型('table'/'column'/'sort')
|
|
363
|
+
|
|
364
|
+
Returns:
|
|
365
|
+
是否安全
|
|
366
|
+
"""
|
|
367
|
+
# 如果不允许${},直接不安全
|
|
368
|
+
if not self.allow_raw_params:
|
|
369
|
+
return False
|
|
370
|
+
|
|
371
|
+
# 检查参数名白名单
|
|
372
|
+
if param_name in self.raw_param_whitelist:
|
|
373
|
+
return True
|
|
374
|
+
|
|
375
|
+
# 检查参数值模式
|
|
376
|
+
if param_value is None:
|
|
377
|
+
return False
|
|
378
|
+
|
|
379
|
+
if not isinstance(param_value, str):
|
|
380
|
+
param_value = str(param_value)
|
|
381
|
+
|
|
382
|
+
# 检查白名单模式
|
|
383
|
+
for pattern in self.RAW_PARAM_WHITELIST_PATTERNS:
|
|
384
|
+
if pattern.match(param_value):
|
|
385
|
+
# 如果指定了参数类型,进行额外验证
|
|
386
|
+
if param_type == 'table' and self.allowed_tables:
|
|
387
|
+
if param_value.lower() in [t.lower() for t in self.allowed_tables]:
|
|
388
|
+
return True
|
|
389
|
+
logger.warning(f"${{{param_name}}} 参数值不在允许的表名单中: {param_value}")
|
|
390
|
+
return False
|
|
391
|
+
if param_type == 'column' and self.allowed_columns:
|
|
392
|
+
if param_value.lower() in [c.lower() for c in self.allowed_columns]:
|
|
393
|
+
return True
|
|
394
|
+
logger.warning(f"${{{param_name}}} 参数值不在允许的字段名单中: {param_value}")
|
|
395
|
+
return False
|
|
396
|
+
return True
|
|
397
|
+
|
|
398
|
+
logger.warning(f"${{{param_name}}} 参数值不在白名单中: {param_value}")
|
|
399
|
+
return False
|
|
400
|
+
|
|
401
|
+
def validate_table_name(self, table_name: str) -> bool:
|
|
402
|
+
"""
|
|
403
|
+
验证表名是否在允许的白名单中
|
|
404
|
+
|
|
405
|
+
Args:
|
|
406
|
+
table_name: 表名
|
|
407
|
+
|
|
408
|
+
Returns:
|
|
409
|
+
是否允许
|
|
410
|
+
"""
|
|
411
|
+
if not self.allowed_tables:
|
|
412
|
+
return True
|
|
413
|
+
|
|
414
|
+
return table_name.lower() in [t.lower() for t in self.allowed_tables]
|
|
415
|
+
|
|
416
|
+
def validate_column_name(self, column_name: str) -> bool:
|
|
417
|
+
"""
|
|
418
|
+
验证字段名是否在允许的白名单中
|
|
419
|
+
|
|
420
|
+
Args:
|
|
421
|
+
column_name: 字段名
|
|
422
|
+
|
|
423
|
+
Returns:
|
|
424
|
+
是否允许
|
|
425
|
+
"""
|
|
426
|
+
if not self.allowed_columns:
|
|
427
|
+
return True
|
|
428
|
+
|
|
429
|
+
return column_name.lower() in [c.lower() for c in self.allowed_columns]
|
|
430
|
+
|
|
431
|
+
def detect_batch(self, params: Dict[str, Any]) -> Dict[str, SQLInjectionLevel]:
|
|
432
|
+
"""
|
|
433
|
+
批量检测多个参数
|
|
434
|
+
|
|
435
|
+
Args:
|
|
436
|
+
params: 参数字典
|
|
437
|
+
|
|
438
|
+
Returns:
|
|
439
|
+
参数名到风险等级的映射
|
|
440
|
+
"""
|
|
441
|
+
results = {}
|
|
442
|
+
for key, value in params.items():
|
|
443
|
+
results[key] = self.detect(value)
|
|
444
|
+
return results
|
|
445
|
+
|
|
446
|
+
def is_blocked(self, value: Any) -> bool:
|
|
447
|
+
"""
|
|
448
|
+
判断值是否会被阻止(风险级别超过允许的最大值)
|
|
449
|
+
|
|
450
|
+
Args:
|
|
451
|
+
value: 待检测的值
|
|
452
|
+
|
|
453
|
+
Returns:
|
|
454
|
+
是否被阻止
|
|
455
|
+
"""
|
|
456
|
+
level = self.detect(value)
|
|
457
|
+
return level.value > self.max_risk_level.value
|
|
458
|
+
|
|
459
|
+
def is_safe(self, value: Any) -> bool:
|
|
460
|
+
"""
|
|
461
|
+
判断值是否安全
|
|
462
|
+
|
|
463
|
+
Args:
|
|
464
|
+
value: 待检测的值
|
|
465
|
+
|
|
466
|
+
Returns:
|
|
467
|
+
是否安全
|
|
468
|
+
"""
|
|
469
|
+
return not self.is_blocked(value)
|
|
470
|
+
|
|
471
|
+
def sanitize(self, value: Any) -> Any:
|
|
472
|
+
"""
|
|
473
|
+
清理值,移除可能的注入内容(修复版:只移除注释和特殊字符,不移除关键字)
|
|
474
|
+
|
|
475
|
+
Args:
|
|
476
|
+
value: 待清理的值
|
|
477
|
+
|
|
478
|
+
Returns:
|
|
479
|
+
清理后的值
|
|
480
|
+
"""
|
|
481
|
+
if not self.enabled:
|
|
482
|
+
return value
|
|
483
|
+
|
|
484
|
+
if value is None:
|
|
485
|
+
return None
|
|
486
|
+
|
|
487
|
+
if not isinstance(value, str):
|
|
488
|
+
return value
|
|
489
|
+
|
|
490
|
+
# 移除SQL注释(防止注释绕过)
|
|
491
|
+
value = re.sub(r'--.*$', '', value, flags=re.MULTILINE)
|
|
492
|
+
value = re.sub(r'#.*$', '', value, flags=re.MULTILINE)
|
|
493
|
+
value = re.sub(r'/\*.*?\*/', '', value, flags=re.DOTALL)
|
|
494
|
+
|
|
495
|
+
# 移除多余的分号(防止堆叠查询)
|
|
496
|
+
value = value.replace(';;', ';')
|
|
497
|
+
|
|
498
|
+
return value.strip()
|
|
499
|
+
|
|
500
|
+
def sanitize_sql(self, sql: str) -> str:
|
|
501
|
+
"""
|
|
502
|
+
清理SQL语句,移除危险内容
|
|
503
|
+
|
|
504
|
+
Args:
|
|
505
|
+
sql: 待清理的SQL
|
|
506
|
+
|
|
507
|
+
Returns:
|
|
508
|
+
清理后的SQL
|
|
509
|
+
"""
|
|
510
|
+
if not self.enabled:
|
|
511
|
+
return sql
|
|
512
|
+
|
|
513
|
+
if sql is None:
|
|
514
|
+
return ''
|
|
515
|
+
|
|
516
|
+
if not isinstance(sql, str):
|
|
517
|
+
sql = str(sql)
|
|
518
|
+
|
|
519
|
+
# 移除SQL注释
|
|
520
|
+
sql = re.sub(r'--.*$', '', sql, flags=re.MULTILINE)
|
|
521
|
+
sql = re.sub(r'#.*$', '', sql, flags=re.MULTILINE)
|
|
522
|
+
sql = re.sub(r'/\*.*?\*/', '', sql, flags=re.DOTALL)
|
|
523
|
+
|
|
524
|
+
# 移除多余分号(防止堆叠查询)
|
|
525
|
+
sql = re.sub(r';+\s*$', ';', sql.strip())
|
|
526
|
+
|
|
527
|
+
return sql
|
|
528
|
+
|
|
529
|
+
def get_detection_details(self, value: Any) -> Dict[str, Any]:
|
|
530
|
+
"""
|
|
531
|
+
获取检测详情
|
|
532
|
+
|
|
533
|
+
Args:
|
|
534
|
+
value: 待检测的值
|
|
535
|
+
|
|
536
|
+
Returns:
|
|
537
|
+
检测详情字典
|
|
538
|
+
"""
|
|
539
|
+
if not self.enabled:
|
|
540
|
+
return {'level': SQLInjectionLevel.NONE, 'patterns': [], 'ast_analysis': 'disabled'}
|
|
541
|
+
|
|
542
|
+
if value is None:
|
|
543
|
+
return {'level': SQLInjectionLevel.NONE, 'patterns': [], 'ast_analysis': 'disabled'}
|
|
544
|
+
|
|
545
|
+
if not isinstance(value, str):
|
|
546
|
+
value = str(value)
|
|
547
|
+
|
|
548
|
+
matched_patterns = []
|
|
549
|
+
max_level = SQLInjectionLevel.NONE
|
|
550
|
+
|
|
551
|
+
for pattern in self.INJECTION_PATTERNS:
|
|
552
|
+
match = pattern.pattern.search(value)
|
|
553
|
+
if match:
|
|
554
|
+
matched_patterns.append({
|
|
555
|
+
'pattern': pattern.pattern.pattern,
|
|
556
|
+
'description': pattern.description,
|
|
557
|
+
'level': pattern.level.name,
|
|
558
|
+
'match': match.group(0)
|
|
559
|
+
})
|
|
560
|
+
if pattern.level.value > max_level.value:
|
|
561
|
+
max_level = pattern.level
|
|
562
|
+
|
|
563
|
+
# AST分析结果
|
|
564
|
+
ast_analysis = 'disabled'
|
|
565
|
+
if self.enable_ast_validation:
|
|
566
|
+
try:
|
|
567
|
+
sqlglot = self._load_sqlglot()
|
|
568
|
+
if sqlglot is not None:
|
|
569
|
+
parsed = sqlglot.parse_one(value)
|
|
570
|
+
ast_analysis = 'valid'
|
|
571
|
+
else:
|
|
572
|
+
ast_analysis = 'sqlglot_not_available'
|
|
573
|
+
except Exception as e:
|
|
574
|
+
ast_analysis = f'parse_error: {str(e)[:50]}'
|
|
575
|
+
|
|
576
|
+
return {
|
|
577
|
+
'level': max_level.name,
|
|
578
|
+
'patterns': matched_patterns,
|
|
579
|
+
'is_blocked': max_level.value > self.max_risk_level.value,
|
|
580
|
+
'ast_analysis': ast_analysis
|
|
581
|
+
}
|
|
582
|
+
|
|
583
|
+
def validate_sql(self, sql: str) -> Dict[str, Any]:
|
|
584
|
+
"""
|
|
585
|
+
完整验证SQL语句(结合正则和AST)
|
|
586
|
+
|
|
587
|
+
Args:
|
|
588
|
+
sql: SQL语句
|
|
589
|
+
|
|
590
|
+
Returns:
|
|
591
|
+
验证结果
|
|
592
|
+
"""
|
|
593
|
+
result = {
|
|
594
|
+
'is_valid': True,
|
|
595
|
+
'errors': [],
|
|
596
|
+
'warnings': [],
|
|
597
|
+
'ast_valid': None
|
|
598
|
+
}
|
|
599
|
+
|
|
600
|
+
if self.detect_ddl(sql):
|
|
601
|
+
result['is_valid'] = False
|
|
602
|
+
result['errors'].append('DDL语句被阻止')
|
|
603
|
+
|
|
604
|
+
injection_result = self.get_detection_details(sql)
|
|
605
|
+
if injection_result['is_blocked']:
|
|
606
|
+
result['is_valid'] = False
|
|
607
|
+
result['errors'].append(f"SQL注入风险: {injection_result['level']}")
|
|
608
|
+
|
|
609
|
+
if injection_result['level'] != SQLInjectionLevel.NONE:
|
|
610
|
+
result['warnings'].append(f"检测到潜在风险: {injection_result['level']}")
|
|
611
|
+
|
|
612
|
+
result['ast_valid'] = injection_result.get('ast_analysis', None)
|
|
613
|
+
|
|
614
|
+
return result
|
|
615
|
+
|
|
616
|
+
|
|
617
|
+
# 全局默认检测器实例(生产环境严格模式)
|
|
618
|
+
DEFAULT_DETECTOR = SQLInjectionDetector(
|
|
619
|
+
enabled=True,
|
|
620
|
+
max_risk_level=SQLInjectionLevel.LOW,
|
|
621
|
+
block_ddl=True,
|
|
622
|
+
allow_raw_params=False
|
|
623
|
+
)
|
|
624
|
+
|
|
625
|
+
|
|
626
|
+
def check_sql_injection(value: Any) -> bool:
|
|
627
|
+
"""
|
|
628
|
+
便捷函数:检查值是否安全
|
|
629
|
+
|
|
630
|
+
Args:
|
|
631
|
+
value: 待检查的值
|
|
632
|
+
|
|
633
|
+
Returns:
|
|
634
|
+
是否安全
|
|
635
|
+
"""
|
|
636
|
+
return DEFAULT_DETECTOR.is_safe(value)
|
|
637
|
+
|
|
638
|
+
|
|
639
|
+
def sanitize_sql_value(value: Any) -> Any:
|
|
640
|
+
"""
|
|
641
|
+
便捷函数:清理SQL值
|
|
642
|
+
|
|
643
|
+
Args:
|
|
644
|
+
value: 待清理的值
|
|
645
|
+
|
|
646
|
+
Returns:
|
|
647
|
+
清理后的值
|
|
648
|
+
"""
|
|
649
|
+
return DEFAULT_DETECTOR.sanitize(value)
|
|
650
|
+
|
|
651
|
+
|
|
652
|
+
def is_ddl_blocked(sql: str) -> bool:
|
|
653
|
+
"""
|
|
654
|
+
便捷函数:检查DDL语句是否被阻止
|
|
655
|
+
|
|
656
|
+
Args:
|
|
657
|
+
sql: SQL语句
|
|
658
|
+
|
|
659
|
+
Returns:
|
|
660
|
+
是否被阻止
|
|
661
|
+
"""
|
|
662
|
+
return DEFAULT_DETECTOR.is_ddl_blocked(sql)
|
|
663
|
+
|
|
664
|
+
|
|
665
|
+
def validate_sql(sql: str) -> Dict[str, Any]:
|
|
666
|
+
"""
|
|
667
|
+
便捷函数:验证SQL语句
|
|
668
|
+
|
|
669
|
+
Args:
|
|
670
|
+
sql: SQL语句
|
|
671
|
+
|
|
672
|
+
Returns:
|
|
673
|
+
验证结果
|
|
674
|
+
"""
|
|
675
|
+
return DEFAULT_DETECTOR.validate_sql(sql)
|