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.
Files changed (175) hide show
  1. spring/__init__.py +66 -0
  2. spring/ai/__init__.py +78 -0
  3. spring/ai/advisors.py +139 -0
  4. spring/ai/annotations.py +74 -0
  5. spring/ai/autoconfig.py +481 -0
  6. spring/ai/core.py +391 -0
  7. spring/ai/etl.py +188 -0
  8. spring/ai/memory.py +109 -0
  9. spring/ai/observability.py +129 -0
  10. spring/ai/providers.py +789 -0
  11. spring/ai/resilience.py +258 -0
  12. spring/ai/tools.py +106 -0
  13. spring/ai/vectorstore.py +303 -0
  14. spring/annotations/__init__.py +188 -0
  15. spring/annotations/cache.py +126 -0
  16. spring/annotations/cloud.py +207 -0
  17. spring/annotations/conditional.py +272 -0
  18. spring/annotations/core.py +864 -0
  19. spring/annotations/messaging.py +107 -0
  20. spring/aop/__init__.py +4 -0
  21. spring/aop/cloud_aop.py +404 -0
  22. spring/aop/comprehensive_aop.py +1015 -0
  23. spring/aop/method_interceptor.py +19 -0
  24. spring/aop/proxy_factory.py +55 -0
  25. spring/cloud/__init__.py +76 -0
  26. spring/cloud/discovery.py +364 -0
  27. spring/cloud/feign.py +469 -0
  28. spring/cloud/gateway.py +452 -0
  29. spring/cloud/load_balancer.py +149 -0
  30. spring/cloud/seata.py +557 -0
  31. spring/cloud/sentinel.py +525 -0
  32. spring/cloud/tracer.py +337 -0
  33. spring/config/__init__.py +21 -0
  34. spring/config/binding.py +206 -0
  35. spring/config/config_loader.py +405 -0
  36. spring/context/__init__.py +13 -0
  37. spring/context/application_context.py +589 -0
  38. spring/context/bean_definition.py +70 -0
  39. spring/context/bean_factory.py +1052 -0
  40. spring/context/registry.py +58 -0
  41. spring/context/scanner.py +106 -0
  42. spring/core/__init__.py +3 -0
  43. spring/core/graceful_shutdown.py +196 -0
  44. spring/core/typing_utils.py +50 -0
  45. spring/csv/__init__.py +52 -0
  46. spring/csv/annotations.py +402 -0
  47. spring/csv/converters.py +69 -0
  48. spring/csv/easy_csv.py +95 -0
  49. spring/csv/exceptions.py +27 -0
  50. spring/csv/reader.py +195 -0
  51. spring/csv/writer.py +155 -0
  52. spring/data/__init__.py +54 -0
  53. spring/data/page.py +181 -0
  54. spring/data/repository.py +274 -0
  55. spring/data/specification.py +228 -0
  56. spring/datasource/__init__.py +66 -0
  57. spring/datasource/annotations.py +133 -0
  58. spring/datasource/context.py +69 -0
  59. spring/datasource/dynamic.py +148 -0
  60. spring/event/__init__.py +7 -0
  61. spring/event/publisher.py +69 -0
  62. spring/excel/__init__.py +51 -0
  63. spring/excel/annotations.py +405 -0
  64. spring/excel/converters.py +231 -0
  65. spring/excel/easy_excel.py +94 -0
  66. spring/excel/exceptions.py +31 -0
  67. spring/excel/reader.py +254 -0
  68. spring/excel/style.py +95 -0
  69. spring/excel/writer.py +197 -0
  70. spring/i18n/__init__.py +97 -0
  71. spring/i18n/accessor.py +94 -0
  72. spring/i18n/auto_config.py +177 -0
  73. spring/i18n/holder.py +106 -0
  74. spring/i18n/locale.py +152 -0
  75. spring/i18n/locale_resolver.py +367 -0
  76. spring/i18n/message_source.py +250 -0
  77. spring/i18n/middleware.py +79 -0
  78. spring/i18n/properties.py +168 -0
  79. spring/i18n/sources.py +255 -0
  80. spring/logging/__init__.py +1 -0
  81. spring/logging/loguru_logger.py +228 -0
  82. spring/main.py +378 -0
  83. spring/messaging/__init__.py +1 -0
  84. spring/messaging/rabbitmq.py +302 -0
  85. spring/monitoring/__init__.py +1 -0
  86. spring/monitoring/prometheus.py +199 -0
  87. spring/orm/__init__.py +258 -0
  88. spring/orm/database.py +222 -0
  89. spring/orm/ddl_auto.py +1217 -0
  90. spring/orm/migration.py +419 -0
  91. spring/orm/mybatis_integration.py +400 -0
  92. spring/orm/pymybatis/__init__.py +86 -0
  93. spring/orm/pymybatis/annotations/__init__.py +30 -0
  94. spring/orm/pymybatis/annotations/annotations.py +332 -0
  95. spring/orm/pymybatis/cache/__init__.py +47 -0
  96. spring/orm/pymybatis/cache/cache.py +371 -0
  97. spring/orm/pymybatis/cache/redis_cache.py +434 -0
  98. spring/orm/pymybatis/circuit_breaker/__init__.py +21 -0
  99. spring/orm/pymybatis/circuit_breaker/circuit_breaker.py +424 -0
  100. spring/orm/pymybatis/configuration.py +525 -0
  101. spring/orm/pymybatis/core/__init__.py +10 -0
  102. spring/orm/pymybatis/core/sql_session.py +1382 -0
  103. spring/orm/pymybatis/core/sql_session_factory.py +76 -0
  104. spring/orm/pymybatis/dialect/__init__.py +9 -0
  105. spring/orm/pymybatis/dialect/dialect.py +445 -0
  106. spring/orm/pymybatis/dynamic_sql/__init__.py +9 -0
  107. spring/orm/pymybatis/dynamic_sql/dynamic_sql.py +900 -0
  108. spring/orm/pymybatis/interceptor/__init__.py +31 -0
  109. spring/orm/pymybatis/interceptor/interceptor.py +427 -0
  110. spring/orm/pymybatis/mapper/__init__.py +9 -0
  111. spring/orm/pymybatis/mapper/mapper.py +540 -0
  112. spring/orm/pymybatis/metrics/__init__.py +41 -0
  113. spring/orm/pymybatis/metrics/metrics.py +595 -0
  114. spring/orm/pymybatis/pool/__init__.py +9 -0
  115. spring/orm/pymybatis/pool/connection_pool.py +711 -0
  116. spring/orm/pymybatis/security/__init__.py +19 -0
  117. spring/orm/pymybatis/security/access_control.py +415 -0
  118. spring/orm/pymybatis/security/password_encoder.py +293 -0
  119. spring/orm/pymybatis/security/sensitive_data_masker.py +326 -0
  120. spring/orm/pymybatis/security/sql_injection_detector.py +675 -0
  121. spring/orm/pymybatis/transaction/__init__.py +9 -0
  122. spring/orm/pymybatis/transaction/transaction.py +288 -0
  123. spring/orm/pymybatis/type_handler/__init__.py +37 -0
  124. spring/orm/pymybatis/type_handler/type_handler.py +473 -0
  125. spring/orm/pymybatis/version.py +9 -0
  126. spring/orm/pymybatis/xml_parser/__init__.py +9 -0
  127. spring/orm/pymybatis/xml_parser/xml_parser.py +761 -0
  128. spring/retry/__init__.py +12 -0
  129. spring/retry/retry_annotations.py +71 -0
  130. spring/retry/retry_decorator.py +155 -0
  131. spring/scheduling/__init__.py +3 -0
  132. spring/scheduling/scheduler.py +389 -0
  133. spring/security/__init__.py +39 -0
  134. spring/security/jwt_utils.py +281 -0
  135. spring/security/replay_protection.py +206 -0
  136. spring/security/secret_manager.py +226 -0
  137. spring/security/security_aop.py +248 -0
  138. spring/security/security_context.py +172 -0
  139. spring/test/__init__.py +45 -0
  140. spring/test/slicing.py +341 -0
  141. spring/tracing/__init__.py +11 -0
  142. spring/tracing/skywalking.py +229 -0
  143. spring/tx/__init__.py +52 -0
  144. spring/tx/events.py +172 -0
  145. spring/tx/synchronization.py +143 -0
  146. spring/utils/__init__.py +5 -0
  147. spring/utils/banner.py +32 -0
  148. spring/utils/logger.py +73 -0
  149. spring/utils/redis_client.py +526 -0
  150. spring/validation/__init__.py +55 -0
  151. spring/validation/aop.py +141 -0
  152. spring/validation/constraints.py +357 -0
  153. spring/validation/exceptions.py +55 -0
  154. spring/validation/validator.py +139 -0
  155. spring/web/__init__.py +12 -0
  156. spring/web/actuator.py +319 -0
  157. spring/web/exception_handler.py +61 -0
  158. spring/web/health.py +399 -0
  159. spring/web/interceptor.py +91 -0
  160. spring/web/result.py +44 -0
  161. spring/web/swagger.py +601 -0
  162. spring/web/web_context.py +755 -0
  163. spring/websocket/__init__.py +86 -0
  164. spring/websocket/annotations.py +169 -0
  165. spring/websocket/broker.py +238 -0
  166. spring/websocket/exceptions.py +26 -0
  167. spring/websocket/handler.py +243 -0
  168. spring/websocket/router.py +526 -0
  169. spring/websocket/session.py +216 -0
  170. springbootai-1.8.0.dist-info/METADATA +2796 -0
  171. springbootai-1.8.0.dist-info/RECORD +175 -0
  172. springbootai-1.8.0.dist-info/WHEEL +5 -0
  173. springbootai-1.8.0.dist-info/entry_points.txt +2 -0
  174. springbootai-1.8.0.dist-info/licenses/LICENSE +7 -0
  175. 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)
@@ -0,0 +1,9 @@
1
+ """
2
+ PyMyBatis事务管理模块
3
+
4
+ 实现事务隔离级别控制、事务边界管理
5
+ """
6
+
7
+ from .transaction import Transaction, TransactionManager, TransactionIsolationLevel
8
+
9
+ __all__ = ['Transaction', 'TransactionManager', 'TransactionIsolationLevel']