Skip to main content

sz_orm_core/
guard.rs

1//! 防全表 UPDATE/DELETE 攻击守卫(Safe SQL Guard)
2//!
3//! 对应文档 6.8 节改进项 26(防全表 UPDATE/DELETE 攻击拦截)。
4//!
5//! # 核心概念
6//!
7//! - **GuardPolicy**:守卫策略,配置是否允许无 WHERE 子句的 UPDATE/DELETE
8//! - **SafeSqlGuard**:守卫实例,对 SQL 进行检查并拦截危险操作
9//! - **GuardError**:守卫错误类型
10//!
11//! # 设计灵感
12//!
13//! - MyBatis-Plus `block-attack-inner-interceptor`(阻止全表更新/删除)
14//! - MySQL `--safe-updates` 模式
15//! - Hibernate `hibernate.query.mutation_strategy`
16//!
17//! # 使用示例
18//!
19//! ```no_run
20//! use sz_orm_core::guard::{SafeSqlGuard, GuardPolicy};
21//!
22//! let guard = SafeSqlGuard::new(GuardPolicy::Strict);
23//! // 拦截无 WHERE 子句的 UPDATE
24//! assert!(guard.check("UPDATE users SET name = 'a'").is_err());
25//! // 拦截无 WHERE 子句的 DELETE
26//! assert!(guard.check("DELETE FROM users").is_err());
27//! // 允许带 WHERE 子句的 UPDATE
28//! assert!(guard.check("UPDATE users SET name = 'a' WHERE id = 1").is_ok());
29//! ```
30
31// ============================================================================
32// GuardError — 守卫错误类型
33// ============================================================================
34
35/// 守卫错误类型
36#[derive(Debug)]
37pub enum GuardError {
38    /// 全表 UPDATE(无 WHERE 子句)
39    FullTableUpdate {
40        /// 表名
41        table: String,
42    },
43    /// 全表 DELETE(无 WHERE 子句)
44    FullTableDelete {
45        /// 表名
46        table: String,
47    },
48    /// SQL 解析失败
49    ParseError(String),
50}
51
52impl std::fmt::Display for GuardError {
53    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
54        match self {
55            GuardError::FullTableUpdate { table } => write!(
56                f,
57                "Blocked full-table UPDATE on `{}` (no WHERE clause). Add WHERE clause or use GuardPolicy::Permissive.",
58                table
59            ),
60            GuardError::FullTableDelete { table } => write!(
61                f,
62                "Blocked full-table DELETE on `{}` (no WHERE clause). Add WHERE clause or use GuardPolicy::Permissive.",
63                table
64            ),
65            GuardError::ParseError(msg) => write!(f, "SQL parse error in guard: {}", msg),
66        }
67    }
68}
69
70impl std::error::Error for GuardError {}
71
72/// 守卫结果
73pub type GuardResult<T> = Result<T, GuardError>;
74
75// ============================================================================
76// GuardPolicy — 守卫策略
77// ============================================================================
78
79/// 守卫策略
80#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
81pub enum GuardPolicy {
82    /// 严格模式(默认):禁止任何无 WHERE 子句的 UPDATE/DELETE
83    #[default]
84    Strict,
85    /// 宽松模式:允许无 WHERE 子句的 UPDATE/DELETE(仅记录日志,不拦截)
86    Permissive,
87    /// 完全关闭守卫
88    Disabled,
89}
90
91// ============================================================================
92// SafeSqlGuard — 守卫实例
93// ============================================================================
94
95/// 防全表 UPDATE/DELETE 攻击守卫
96///
97/// 通过 SQL 解析检测是否存在 WHERE 子句,对全表 UPDATE/DELETE 操作进行拦截。
98///
99/// # 示例
100///
101/// ```
102/// use sz_orm_core::guard::{SafeSqlGuard, GuardPolicy, GuardError};
103///
104/// let guard = SafeSqlGuard::new(GuardPolicy::Strict);
105///
106/// // 拦截无 WHERE 子句的 UPDATE
107/// assert!(matches!(
108///     guard.check("UPDATE users SET name = 'a'"),
109///     Err(GuardError::FullTableUpdate { .. })
110/// ));
111///
112/// // 允许带 WHERE 子句的 UPDATE
113/// assert!(guard.check("UPDATE users SET name = 'a' WHERE id = 1").is_ok());
114/// ```
115#[derive(Debug, Clone, Copy)]
116pub struct SafeSqlGuard {
117    /// 策略
118    pub policy: GuardPolicy,
119}
120
121impl SafeSqlGuard {
122    /// 创建守卫实例
123    pub fn new(policy: GuardPolicy) -> Self {
124        Self { policy }
125    }
126
127    /// 创建默认(严格)守卫
128    pub fn strict() -> Self {
129        Self::new(GuardPolicy::Strict)
130    }
131
132    /// 创建宽松守卫
133    pub fn permissive() -> Self {
134        Self::new(GuardPolicy::Permissive)
135    }
136
137    /// 创建关闭守卫
138    pub fn disabled() -> Self {
139        Self::new(GuardPolicy::Disabled)
140    }
141
142    /// 检查 SQL 是否安全(拦截全表 UPDATE/DELETE)
143    pub fn check(&self, sql: &str) -> GuardResult<()> {
144        match self.policy {
145            GuardPolicy::Disabled => Ok(()),
146            GuardPolicy::Permissive => {
147                // 宽松模式:仅记录但不拦截,简单返回 Ok
148                // 生产环境可在此处接入日志系统
149                Ok(())
150            }
151            GuardPolicy::Strict => check_sql_strict(sql),
152        }
153    }
154
155    /// 检查 SQL 是否安全(与 check 等价,但明确表示检查 UPDATE)
156    pub fn check_update(&self, sql: &str) -> GuardResult<()> {
157        self.check(sql)
158    }
159
160    /// 检查 SQL 是否安全(与 check 等价,但明确表示检查 DELETE)
161    pub fn check_delete(&self, sql: &str) -> GuardResult<()> {
162        self.check(sql)
163    }
164}
165
166impl Default for SafeSqlGuard {
167    fn default() -> Self {
168        Self::strict()
169    }
170}
171
172// ============================================================================
173// 内部解析函数
174// ============================================================================
175
176/// 严格模式检查 SQL:检测无 WHERE 子句的 UPDATE/DELETE
177///
178/// # 已知局限
179///
180/// 本守卫使用简化的字符串匹配,**不能**完全防御以下绕过场景:
181/// - 子查询中的 WHERE 被误认为外层 UPDATE/DELETE 的 WHERE
182///   (如 `UPDATE users SET name = (SELECT name FROM other WHERE id = 1)`)
183/// - 字符串字面量中的 WHERE 关键字
184///
185/// 完全防御需要 SQL 解析器(如 `sqlparser` crate)。
186/// 本守卫的目标是拦截**明显的**全表 UPDATE/DELETE(无 WHERE 子句)。
187fn check_sql_strict(sql: &str) -> GuardResult<()> {
188    let normalized = normalize_sql(sql);
189    let upper = normalized.to_uppercase();
190
191    // 检测是否为 UPDATE 语句(UPDATE 关键字开头)
192    if upper.starts_with("UPDATE") {
193        let table = extract_update_table(&normalized, &upper);
194        // WHERE 必须在 SET 之后(避免 SET 之前的 WHERE 被误判)
195        if !has_where_after_keyword(&upper, "SET") {
196            return Err(GuardError::FullTableUpdate {
197                table: table.unwrap_or_else(|| "unknown".to_string()),
198            });
199        }
200    }
201
202    // 检测是否为 DELETE 语句
203    if upper.starts_with("DELETE") {
204        let table = extract_delete_table(&normalized, &upper);
205        // WHERE 必须在 FROM 之后
206        if !has_where_after_keyword(&upper, "FROM") {
207            return Err(GuardError::FullTableDelete {
208                table: table.unwrap_or_else(|| "unknown".to_string()),
209            });
210        }
211    }
212
213    Ok(())
214}
215
216/// 规范化 SQL:去除多余空白、注释、换行
217fn normalize_sql(sql: &str) -> String {
218    // 去除单行注释(-- ...)
219    let without_line_comments: String = sql
220        .lines()
221        .map(|line| {
222            if let Some(idx) = line.find("--") {
223                &line[..idx]
224            } else {
225                line
226            }
227        })
228        .collect::<Vec<&str>>()
229        .join(" ");
230
231    // 去除多行注释(/* ... */)
232    let mut without_block_comments = String::with_capacity(without_line_comments.len());
233    let mut in_block_comment = false;
234    let mut chars = without_line_comments.chars().peekable();
235    while let Some(c) = chars.next() {
236        if c == '/' && chars.peek() == Some(&'*') {
237            in_block_comment = true;
238            chars.next(); // 消费 '*'
239            continue;
240        }
241        if in_block_comment && c == '*' && chars.peek() == Some(&'/') {
242            in_block_comment = false;
243            chars.next(); // 消费 '/'
244            continue;
245        }
246        if !in_block_comment {
247            without_block_comments.push(c);
248        }
249    }
250
251    // 折叠空白
252    let mut result = String::with_capacity(without_block_comments.len());
253    let mut prev_whitespace = false;
254    for c in without_block_comments.chars() {
255        if c.is_whitespace() {
256            if !prev_whitespace {
257                result.push(' ');
258                prev_whitespace = true;
259            }
260        } else {
261            result.push(c);
262            prev_whitespace = false;
263        }
264    }
265    result.trim().to_string()
266}
267
268/// 检测 SQL 在指定关键字(如 "SET" 或 "FROM")之后是否包含 WHERE 子句
269///
270/// 规则:
271/// - 必须找到 `keyword` 关键字(作为独立词)
272/// - 在 `keyword` 之后必须包含 `WHERE` 关键字(独立词,且**括号深度为 0**)
273/// - WHERE 子句后必须有非空内容
274///
275/// 此函数防御以下绕过:
276/// - "SET 之前的 WHERE 被误认为 UPDATE 的 WHERE"
277/// - "子查询中的 WHERE 被误认为外层 UPDATE/DELETE 的 WHERE"
278///
279/// 通过括号深度跟踪识别子查询:只有 depth=0 时的 WHERE 才算外层 WHERE。
280fn has_where_after_keyword(upper_sql: &str, keyword: &str) -> bool {
281    // 找到 keyword 的位置(首次出现,独立词匹配)
282    let kw_pos = match find_keyword_independent(upper_sql, keyword) {
283        Some(idx) => idx,
284        None => return false,
285    };
286
287    // 在 keyword 之后查找 WHERE 关键字(独立词,且括号深度为 0)
288    let after_kw = &upper_sql[kw_pos + keyword.len()..];
289    let where_idx = match find_where_at_depth_zero(after_kw) {
290        Some(idx) => idx,
291        None => return false,
292    };
293
294    // WHERE 之后必须有非空内容
295    let after_where = after_kw[where_idx + 5..].trim();
296    !after_where.is_empty()
297}
298
299/// 在 SQL 中查找 WHERE 关键字的位置(独立词,且括号深度为 0)
300///
301/// 扫描字符串,跟踪括号深度(`(` +1, `)` -1),仅返回深度为 0 时的 WHERE 位置。
302/// 这样可以正确区分外层 WHERE 和子查询中的 WHERE。
303///
304/// 注意:这是简化的实现,不处理字符串字面量中的括号(如 `WHERE name = '('`)。
305/// 完全处理需要 SQL 词法分析器。
306fn find_where_at_depth_zero(sql: &str) -> Option<usize> {
307    let bytes = sql.as_bytes();
308    let mut depth: i32 = 0;
309    let mut i = 0;
310
311    while i + 5 <= bytes.len() {
312        let c = bytes[i];
313
314        // 跟踪括号深度
315        if c == b'(' {
316            depth += 1;
317            i += 1;
318            continue;
319        }
320        if c == b')' {
321            depth -= 1;
322            if depth < 0 {
323                depth = 0;
324            }
325            i += 1;
326            continue;
327        }
328
329        // 仅在深度为 0 时查找 WHERE
330        if depth == 0 && &bytes[i..i + 5] == b"WHERE" {
331            // 检查前一个字符是否为单词边界
332            let prev_ok = i == 0 || !bytes[i - 1].is_ascii_alphanumeric() && bytes[i - 1] != b'_';
333            // 检查后一个字符是否为单词边界
334            let next_idx = i + 5;
335            let next_ok = next_idx >= bytes.len()
336                || !bytes[next_idx].is_ascii_alphanumeric() && bytes[next_idx] != b'_';
337            if prev_ok && next_ok {
338                return Some(i);
339            }
340        }
341
342        i += 1;
343    }
344
345    None
346}
347
348/// 在 SQL 中查找指定关键字的位置(独立词匹配,大小写敏感——已 to_uppercase)
349///
350/// 关键字必须是纯单词(如 "SET"、"FROM"、"WHERE"),**不应包含空格**。
351/// 函数会检查关键字前后是否为单词边界(非字母数字下划线)。
352fn find_keyword_independent(sql: &str, keyword: &str) -> Option<usize> {
353    let kw_len = keyword.len();
354    if kw_len == 0 || sql.len() < kw_len {
355        return None;
356    }
357
358    let bytes = sql.as_bytes();
359    let kw_bytes = keyword.as_bytes();
360
361    let mut i = 0;
362    while i + kw_len <= bytes.len() {
363        if &bytes[i..i + kw_len] == kw_bytes {
364            // 检查前一个字符是否为单词边界
365            let prev_ok = i == 0 || !bytes[i - 1].is_ascii_alphanumeric() && bytes[i - 1] != b'_';
366            // 检查后一个字符是否为单词边界
367            let next_idx = i + kw_len;
368            let next_ok = next_idx >= bytes.len()
369                || !bytes[next_idx].is_ascii_alphanumeric() && bytes[next_idx] != b'_';
370            if prev_ok && next_ok {
371                return Some(i);
372            }
373        }
374        i += 1;
375    }
376    None
377}
378
379/// 从 UPDATE 语句中提取表名
380///
381/// 支持格式:
382/// - `UPDATE table SET ...`
383/// - `UPDATE schema.table SET ...`
384/// - `UPDATE `table` SET ...`
385/// - `UPDATE "table" SET ...`
386fn extract_update_table(sql: &str, upper: &str) -> Option<String> {
387    // 找到 UPDATE 之后、SET 之前的内容
388    let update_end = upper.find("UPDATE").map(|i| i + 6)?;
389    let set_idx = upper.find(" SET ")?;
390
391    let between = sql[update_end..set_idx].trim();
392
393    // 处理反引号/双引号引用的表名
394    let table = if (between.starts_with('`') && between.ends_with('`'))
395        || (between.starts_with('"') && between.ends_with('"'))
396    {
397        between[1..between.len() - 1].to_string()
398    } else {
399        between.to_string()
400    };
401
402    // 处理 schema.table 格式,只取 table 部分
403    let table = table.rsplit('.').next().unwrap_or(&table).to_string();
404
405    if table.is_empty() {
406        None
407    } else {
408        Some(table)
409    }
410}
411
412/// 从 DELETE 语句中提取表名
413///
414/// 支持格式:
415/// - `DELETE FROM table`
416/// - `DELETE FROM schema.table`
417/// - `DELETE FROM `table``
418/// - `DELETE table WHERE ...`(MySQL 扩展语法)
419fn extract_delete_table(sql: &str, upper: &str) -> Option<String> {
420    // 优先匹配 DELETE FROM table
421    if let Some(from_idx) = upper.find("DELETE FROM") {
422        let after_from = sql[from_idx + 11..].trim_start();
423        // 截取到下一个空格或末尾
424        let end = after_from
425            .find(|c: char| c.is_whitespace())
426            .unwrap_or(after_from.len());
427        let raw = after_from[..end].trim();
428
429        return parse_table_name(raw);
430    }
431
432    // 处理 `DELETE table WHERE ...` 扩展语法
433    if upper.starts_with("DELETE ") {
434        let after_delete = sql[7..].trim_start();
435        let end = after_delete
436            .find(|c: char| c.is_whitespace())
437            .unwrap_or(after_delete.len());
438        let raw = after_delete[..end].trim();
439        return parse_table_name(raw);
440    }
441
442    None
443}
444
445/// 解析表名(去除引号、schema 前缀)
446fn parse_table_name(raw: &str) -> Option<String> {
447    let table = if (raw.starts_with('`') && raw.ends_with('`') && raw.len() >= 2)
448        || (raw.starts_with('"') && raw.ends_with('"') && raw.len() >= 2)
449    {
450        raw[1..raw.len() - 1].to_string()
451    } else {
452        raw.to_string()
453    };
454
455    let table = table.rsplit('.').next().unwrap_or(&table).to_string();
456
457    if table.is_empty() {
458        None
459    } else {
460        Some(table)
461    }
462}
463
464// ============================================================================
465// 单元测试
466// ============================================================================
467
468#[cfg(test)]
469mod tests {
470    use super::*;
471
472    // ===== Strict 模式 - UPDATE 拦截 =====
473
474    #[test]
475    fn test_strict_blocks_update_without_where() {
476        let guard = SafeSqlGuard::strict();
477        let result = guard.check("UPDATE users SET name = 'a'");
478        assert!(matches!(
479            result,
480            Err(GuardError::FullTableUpdate { table }) if table == "users"
481        ));
482    }
483
484    #[test]
485    fn test_strict_blocks_update_without_where_multiple_columns() {
486        let guard = SafeSqlGuard::strict();
487        let result = guard.check("UPDATE users SET name = 'a', age = 30, status = 'active'");
488        assert!(matches!(result, Err(GuardError::FullTableUpdate { .. })));
489    }
490
491    #[test]
492    fn test_strict_allows_update_with_where() {
493        let guard = SafeSqlGuard::strict();
494        let result = guard.check("UPDATE users SET name = 'a' WHERE id = 1");
495        assert!(result.is_ok());
496    }
497
498    #[test]
499    fn test_strict_allows_update_with_where_in() {
500        let guard = SafeSqlGuard::strict();
501        let result = guard.check("UPDATE users SET name = 'a' WHERE id IN (1, 2, 3)");
502        assert!(result.is_ok());
503    }
504
505    #[test]
506    fn test_strict_blocks_update_with_quoted_table() {
507        let guard = SafeSqlGuard::strict();
508        let result = guard.check("UPDATE `users` SET name = 'a'");
509        assert!(matches!(
510            result,
511            Err(GuardError::FullTableUpdate { table }) if table == "users"
512        ));
513    }
514
515    #[test]
516    fn test_strict_blocks_update_with_double_quoted_table() {
517        let guard = SafeSqlGuard::strict();
518        let result = guard.check("UPDATE \"users\" SET name = 'a'");
519        assert!(matches!(
520            result,
521            Err(GuardError::FullTableUpdate { table }) if table == "users"
522        ));
523    }
524
525    #[test]
526    fn test_strict_blocks_update_with_schema_qualified_table() {
527        let guard = SafeSqlGuard::strict();
528        let result = guard.check("UPDATE public.users SET name = 'a'");
529        assert!(matches!(
530            result,
531            Err(GuardError::FullTableUpdate { table }) if table == "users"
532        ));
533    }
534
535    // ===== Strict 模式 - DELETE 拦截 =====
536
537    #[test]
538    fn test_strict_blocks_delete_without_where() {
539        let guard = SafeSqlGuard::strict();
540        let result = guard.check("DELETE FROM users");
541        assert!(matches!(
542            result,
543            Err(GuardError::FullTableDelete { table }) if table == "users"
544        ));
545    }
546
547    #[test]
548    fn test_strict_allows_delete_with_where() {
549        let guard = SafeSqlGuard::strict();
550        let result = guard.check("DELETE FROM users WHERE id = 1");
551        assert!(result.is_ok());
552    }
553
554    #[test]
555    fn test_strict_blocks_delete_with_quoted_table() {
556        let guard = SafeSqlGuard::strict();
557        let result = guard.check("DELETE FROM `users`");
558        assert!(matches!(
559            result,
560            Err(GuardError::FullTableDelete { table }) if table == "users"
561        ));
562    }
563
564    #[test]
565    fn test_strict_blocks_delete_mysql_extension() {
566        // MySQL 扩展语法:DELETE table WHERE ...
567        // 这里测试无 WHERE 的情况
568        let guard = SafeSqlGuard::strict();
569        let result = guard.check("DELETE users");
570        assert!(matches!(
571            result,
572            Err(GuardError::FullTableDelete { table }) if table == "users"
573        ));
574    }
575
576    // ===== Permissive 模式 =====
577
578    #[test]
579    fn test_permissive_allows_update_without_where() {
580        let guard = SafeSqlGuard::permissive();
581        let result = guard.check("UPDATE users SET name = 'a'");
582        assert!(result.is_ok());
583    }
584
585    #[test]
586    fn test_permissive_allows_delete_without_where() {
587        let guard = SafeSqlGuard::permissive();
588        let result = guard.check("DELETE FROM users");
589        assert!(result.is_ok());
590    }
591
592    // ===== Disabled 模式 =====
593
594    #[test]
595    fn test_disabled_allows_everything() {
596        let guard = SafeSqlGuard::disabled();
597        assert!(guard.check("UPDATE users SET name = 'a'").is_ok());
598        assert!(guard.check("DELETE FROM users").is_ok());
599    }
600
601    // ===== 非 UPDATE/DELETE 语句 =====
602
603    #[test]
604    fn test_strict_allows_select_without_where() {
605        let guard = SafeSqlGuard::strict();
606        let result = guard.check("SELECT * FROM users");
607        assert!(result.is_ok());
608    }
609
610    #[test]
611    fn test_strict_allows_insert_without_where() {
612        let guard = SafeSqlGuard::strict();
613        let result = guard.check("INSERT INTO users (name) VALUES ('a')");
614        assert!(result.is_ok());
615    }
616
617    #[test]
618    fn test_strict_allows_create_table() {
619        let guard = SafeSqlGuard::strict();
620        let result = guard.check("CREATE TABLE users (id INT)");
621        assert!(result.is_ok());
622    }
623
624    // ===== 多行/带注释 SQL =====
625
626    #[test]
627    fn test_strict_blocks_multiline_update_without_where() {
628        let guard = SafeSqlGuard::strict();
629        let sql = "UPDATE users\nSET name = 'a',\n    age = 30";
630        let result = guard.check(sql);
631        assert!(matches!(result, Err(GuardError::FullTableUpdate { .. })));
632    }
633
634    #[test]
635    fn test_strict_allows_multiline_update_with_where() {
636        let guard = SafeSqlGuard::strict();
637        let sql = "UPDATE users\nSET name = 'a'\nWHERE id = 1";
638        let result = guard.check(sql);
639        assert!(result.is_ok());
640    }
641
642    #[test]
643    fn test_strict_blocks_update_with_line_comment_only() {
644        let guard = SafeSqlGuard::strict();
645        let sql = "UPDATE users SET name = 'a' -- WHERE id = 1";
646        let result = guard.check(sql);
647        // 注释中的 WHERE 不应被识别为真正的 WHERE 子句
648        assert!(matches!(result, Err(GuardError::FullTableUpdate { .. })));
649    }
650
651    #[test]
652    fn test_strict_blocks_update_with_block_comment_only() {
653        let guard = SafeSqlGuard::strict();
654        let sql = "UPDATE users SET name = 'a' /* WHERE id = 1 */";
655        let result = guard.check(sql);
656        assert!(matches!(result, Err(GuardError::FullTableUpdate { .. })));
657    }
658
659    #[test]
660    fn test_strict_allows_update_with_real_where_and_comment() {
661        let guard = SafeSqlGuard::strict();
662        let sql = "UPDATE users SET name = 'a' /* update name */ WHERE id = 1";
663        let result = guard.check(sql);
664        assert!(result.is_ok());
665    }
666
667    // ===== WHERE 关键字边界检测 =====
668
669    #[test]
670    fn test_strict_does_not_treat_nowhere_as_where() {
671        // WHERE 是子串的情况:比如字段名包含 WHERE
672        let guard = SafeSqlGuard::strict();
673        // 表名/字段名中含 WHERE 子串,不应被误判为 WHERE 子句
674        let sql = "UPDATE my_table SET somewhere = 'x'";
675        let result = guard.check(sql);
676        assert!(matches!(result, Err(GuardError::FullTableUpdate { .. })));
677    }
678
679    // ===== check_update / check_delete 便捷方法 =====
680
681    #[test]
682    fn test_check_update_method() {
683        let guard = SafeSqlGuard::strict();
684        assert!(guard.check_update("UPDATE users SET name = 'a'").is_err());
685        assert!(guard
686            .check_update("UPDATE users SET name = 'a' WHERE id = 1")
687            .is_ok());
688    }
689
690    #[test]
691    fn test_check_delete_method() {
692        let guard = SafeSqlGuard::strict();
693        assert!(guard.check_delete("DELETE FROM users").is_err());
694        assert!(guard.check_delete("DELETE FROM users WHERE id = 1").is_ok());
695    }
696
697    // ===== Default =====
698
699    #[test]
700    fn test_default_guard_is_strict() {
701        let guard = SafeSqlGuard::default();
702        assert_eq!(guard.policy, GuardPolicy::Strict);
703        assert!(guard.check("UPDATE users SET name = 'a'").is_err());
704    }
705
706    // ===== GuardError Display =====
707
708    #[test]
709    fn test_guard_error_display_full_table_update() {
710        let e = GuardError::FullTableUpdate {
711            table: "users".to_string(),
712        };
713        let s = format!("{}", e);
714        assert!(s.contains("Blocked full-table UPDATE"));
715        assert!(s.contains("users"));
716    }
717
718    #[test]
719    fn test_guard_error_display_full_table_delete() {
720        let e = GuardError::FullTableDelete {
721            table: "orders".to_string(),
722        };
723        let s = format!("{}", e);
724        assert!(s.contains("Blocked full-table DELETE"));
725        assert!(s.contains("orders"));
726    }
727
728    // ===== GuardPolicy Default =====
729
730    #[test]
731    fn test_guard_policy_default_is_strict() {
732        let p = GuardPolicy::default();
733        assert_eq!(p, GuardPolicy::Strict);
734    }
735
736    // ===== 边界情况 =====
737
738    #[test]
739    fn test_strict_blocks_truncate() {
740        // TRUNCATE 也是全表删除,但目前未实现 TRUNCATE 检测
741        // 这里仅验证 TRUNCATE 不会被错误识别为 UPDATE/DELETE
742        let guard = SafeSqlGuard::strict();
743        // TRUNCATE 应该是允许的(不在守卫范围)
744        let result = guard.check("TRUNCATE TABLE users");
745        assert!(result.is_ok());
746    }
747
748    #[test]
749    fn test_strict_blocks_update_with_lowercase_keywords() {
750        let guard = SafeSqlGuard::strict();
751        let result = guard.check("update users set name = 'a'");
752        assert!(matches!(result, Err(GuardError::FullTableUpdate { .. })));
753    }
754
755    #[test]
756    fn test_strict_blocks_delete_with_lowercase_keywords() {
757        let guard = SafeSqlGuard::strict();
758        let result = guard.check("delete from users");
759        assert!(matches!(result, Err(GuardError::FullTableDelete { .. })));
760    }
761
762    #[test]
763    fn test_strict_allows_update_with_lowercase_where() {
764        let guard = SafeSqlGuard::strict();
765        let result = guard.check("update users set name = 'a' where id = 1");
766        assert!(result.is_ok());
767    }
768
769    #[test]
770    fn test_strict_blocks_update_with_empty_where() {
771        let guard = SafeSqlGuard::strict();
772        // "UPDATE users SET name = 'a' WHERE" 后面没有内容
773        let result = guard.check("UPDATE users SET name = 'a' WHERE ");
774        assert!(matches!(result, Err(GuardError::FullTableUpdate { .. })));
775    }
776
777    #[test]
778    fn test_normalize_sql_collapses_whitespace() {
779        let sql = "UPDATE   users\n\nSET    name = 'a'";
780        let normalized = normalize_sql(sql);
781        assert_eq!(normalized, "UPDATE users SET name = 'a'");
782    }
783
784    #[test]
785    fn test_normalize_sql_removes_line_comments() {
786        let sql = "UPDATE users -- this is a comment\nSET name = 'a'";
787        let normalized = normalize_sql(sql);
788        assert!(normalized.contains("UPDATE users"));
789        assert!(!normalized.contains("this is a comment"));
790    }
791
792    #[test]
793    fn test_normalize_sql_removes_block_comments() {
794        let sql = "UPDATE users /* block comment */ SET name = 'a'";
795        let normalized = normalize_sql(sql);
796        assert!(!normalized.contains("block comment"));
797        assert!(normalized.contains("UPDATE users"));
798        assert!(normalized.contains("SET name = 'a'"));
799    }
800
801    // ===== 子查询 WHERE 绕过防御(C2 修复回归测试) =====
802
803    #[test]
804    fn test_blocks_update_with_where_only_in_subquery() {
805        // 子查询中的 WHERE 不应被误判为外层 UPDATE 的 WHERE
806        // 这是 C2 修复的核心回归测试
807        let guard = SafeSqlGuard::strict();
808        let sql = "UPDATE users SET name = (SELECT name FROM other WHERE id = 1)";
809        let result = guard.check(sql);
810        assert!(
811            matches!(result, Err(GuardError::FullTableUpdate { table }) if table == "users"),
812            "子查询中的 WHERE 不应被误认为外层 UPDATE 的 WHERE,应拦截"
813        );
814    }
815
816    #[test]
817    fn test_blocks_delete_with_where_only_in_subquery() {
818        // 子查询中的 WHERE 不应被误判为外层 DELETE 的 WHERE
819        let guard = SafeSqlGuard::strict();
820        let sql = "DELETE FROM users WHERE id IN (SELECT id FROM other)";
821        // 上面的 WHERE 是真正的外层 WHERE,应该通过
822        let result = guard.check(sql);
823        assert!(result.is_ok(), "外层 WHERE 应被识别");
824
825        // 但这个 SQL 外层没有 WHERE,子查询里有 WHERE,应该被拦截
826        let sql2 = "DELETE FROM users RETURNING (SELECT id FROM other WHERE x = 1)";
827        let result2 = guard.check(sql2);
828        assert!(
829            matches!(result2, Err(GuardError::FullTableDelete { table }) if table == "users"),
830            "子查询中的 WHERE 不应被误认为外层 DELETE 的 WHERE,应拦截"
831        );
832    }
833
834    #[test]
835    fn test_allows_update_with_real_where_and_subquery_where() {
836        // 外层有 WHERE + 子查询也有 WHERE,应该通过
837        let guard = SafeSqlGuard::strict();
838        let sql = "UPDATE users SET name = 'a' WHERE id IN (SELECT id FROM other WHERE active = 1)";
839        let result = guard.check(sql);
840        assert!(result.is_ok());
841    }
842
843    #[test]
844    fn test_blocks_update_with_field_named_where() {
845        // 字段名包含 WHERE 子串,不应被误判
846        let guard = SafeSqlGuard::strict();
847        let sql = "UPDATE my_table SET somewhere = 'x'";
848        let result = guard.check(sql);
849        assert!(matches!(result, Err(GuardError::FullTableUpdate { .. })));
850    }
851}