Skip to main content

sz_orm_sql_validator/
firewall.rs

1//! SQL 防火墙:运行时拦截危险 SQL
2
3use std::sync::RwLock;
4
5/// SQL 防火墙规则
6#[derive(Debug, Clone)]
7pub struct FirewallRule {
8    /// 规则名称
9    pub name: String,
10    /// 匹配模式(正则表达式)
11    pub pattern: String,
12    /// 规则动作
13    pub action: FirewallAction,
14    /// 例外模式:若匹配则规则不触发(替代正则 lookahead,Rust regex 不支持 lookahead)
15    pub unless_pattern: Option<String>,
16}
17
18/// 防火墙动作
19#[derive(Debug, Clone, PartialEq, Eq)]
20pub enum FirewallAction {
21    /// 阻断
22    Block,
23    /// 记录但放行
24    Log,
25    /// 需要审批
26    RequireApproval,
27}
28
29/// SQL 防火墙
30pub struct SqlFirewall {
31    rules: RwLock<Vec<FirewallRule>>,
32    blocked_count: std::sync::atomic::AtomicU64,
33    logged_count: std::sync::atomic::AtomicU64,
34}
35
36impl SqlFirewall {
37    pub fn new() -> Self {
38        // 默认规则(使用 vec![] 宏避免 clippy::vec_init_then_push 警告)
39        let rules = vec![
40            FirewallRule {
41                name: "block_drop_table".into(),
42                pattern: r"(?i)\bDROP\s+TABLE\b".into(),
43                action: FirewallAction::Block,
44                unless_pattern: None,
45            },
46            FirewallRule {
47                name: "block_truncate".into(),
48                pattern: r"(?i)\bTRUNCATE\b".into(),
49                action: FirewallAction::Block,
50                unless_pattern: None,
51            },
52            FirewallRule {
53                name: "block_drop_database".into(),
54                pattern: r"(?i)\bDROP\s+DATABASE\b".into(),
55                action: FirewallAction::Block,
56                unless_pattern: None,
57            },
58            FirewallRule {
59                name: "block_drop_schema".into(),
60                pattern: r"(?i)\bDROP\s+SCHEMA\b".into(),
61                action: FirewallAction::Block,
62                unless_pattern: None,
63            },
64            FirewallRule {
65                name: "block_alter_table_drop".into(),
66                pattern: r"(?i)\bALTER\s+TABLE\b.*\bDROP\b".into(),
67                action: FirewallAction::Block,
68                unless_pattern: None,
69            },
70            // DELETE FROM without WHERE — Rust regex 不支持 lookahead,改用 unless_pattern 表达"含 WHERE 则放行"
71            FirewallRule {
72                name: "log_delete_without_where".into(),
73                pattern: r"(?i)\bDELETE\s+FROM\b".into(),
74                action: FirewallAction::Block,
75                unless_pattern: Some(r"(?i)\bWHERE\b".into()),
76            },
77            FirewallRule {
78                name: "block_update_without_where".into(),
79                pattern: r"(?i)\bUPDATE\b.*\bSET\b".into(),
80                action: FirewallAction::Block,
81                unless_pattern: Some(r"(?i)\bWHERE\b".into()),
82            },
83            FirewallRule {
84                name: "block_grant".into(),
85                pattern: r"(?i)\bGRANT\b".into(),
86                action: FirewallAction::Block,
87                unless_pattern: None,
88            },
89            FirewallRule {
90                name: "block_revoke".into(),
91                pattern: r"(?i)\bREVOKE\b".into(),
92                action: FirewallAction::Block,
93                unless_pattern: None,
94            },
95        ];
96
97        Self {
98            rules: RwLock::new(rules),
99            blocked_count: std::sync::atomic::AtomicU64::new(0),
100            logged_count: std::sync::atomic::AtomicU64::new(0),
101        }
102    }
103
104    /// 检查 SQL 是否允许执行
105    pub fn check(&self, sql: &str) -> Result<(), FirewallViolation> {
106        let rules = self.rules.read().expect("Firewall rules lock poisoned");
107        for rule in rules.iter() {
108            if let Ok(regex) = regex::Regex::new(&rule.pattern) {
109                if regex.is_match(sql) {
110                    // 检查例外条件:若 unless_pattern 匹配,则跳过此规则
111                    if let Some(unless_pat) = &rule.unless_pattern {
112                        if let Ok(unless_re) = regex::Regex::new(unless_pat) {
113                            if unless_re.is_match(sql) {
114                                continue;
115                            }
116                        }
117                    }
118                    match rule.action {
119                        FirewallAction::Block => {
120                            self.blocked_count
121                                .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
122                            return Err(FirewallViolation {
123                                rule_name: rule.name.clone(),
124                                sql: sql.to_string(),
125                                action: rule.action.clone(),
126                            });
127                        }
128                        FirewallAction::Log => {
129                            self.logged_count
130                                .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
131                        }
132                        FirewallAction::RequireApproval => {
133                            self.logged_count
134                                .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
135                        }
136                    }
137                }
138            }
139        }
140        Ok(())
141    }
142
143    /// 添加自定义规则
144    pub fn add_rule(&self, rule: FirewallRule) {
145        let mut rules = self.rules.write().expect("Firewall rules lock poisoned");
146        rules.push(rule);
147    }
148
149    /// 获取被阻断次数
150    pub fn blocked_count(&self) -> u64 {
151        self.blocked_count
152            .load(std::sync::atomic::Ordering::Relaxed)
153    }
154
155    /// 获取被记录次数
156    pub fn logged_count(&self) -> u64 {
157        self.logged_count.load(std::sync::atomic::Ordering::Relaxed)
158    }
159}
160
161impl Default for SqlFirewall {
162    fn default() -> Self {
163        Self::new()
164    }
165}
166
167/// 防火墙违规
168#[derive(Debug, Clone)]
169pub struct FirewallViolation {
170    pub rule_name: String,
171    pub sql: String,
172    pub action: FirewallAction,
173}
174
175impl std::fmt::Display for FirewallViolation {
176    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
177        write!(f, "SQL blocked by rule '{}': {}", self.rule_name, self.sql)
178    }
179}
180
181impl std::error::Error for FirewallViolation {}
182
183#[cfg(test)]
184mod tests {
185    use super::*;
186
187    #[test]
188    fn test_allows_simple_select() {
189        let fw = SqlFirewall::new();
190        assert!(fw.check("SELECT id, name FROM users WHERE id = 1").is_ok());
191    }
192
193    #[test]
194    fn test_blocks_drop_table() {
195        let fw = SqlFirewall::new();
196        let result = fw.check("DROP TABLE users");
197        assert!(result.is_err());
198        let err = result.unwrap_err();
199        assert_eq!(err.rule_name, "block_drop_table");
200        assert_eq!(err.action, FirewallAction::Block);
201    }
202
203    #[test]
204    fn test_blocks_truncate() {
205        let fw = SqlFirewall::new();
206        assert!(fw.check("TRUNCATE TABLE users").is_err());
207    }
208
209    #[test]
210    fn test_blocks_drop_database() {
211        let fw = SqlFirewall::new();
212        assert!(fw.check("DROP DATABASE prod").is_err());
213    }
214
215    #[test]
216    fn test_blocks_drop_schema() {
217        let fw = SqlFirewall::new();
218        assert!(fw.check("DROP SCHEMA public").is_err());
219    }
220
221    #[test]
222    fn test_blocks_alter_table_drop() {
223        let fw = SqlFirewall::new();
224        assert!(fw.check("ALTER TABLE users DROP COLUMN password").is_err());
225    }
226
227    #[test]
228    fn test_blocks_delete_without_where() {
229        let fw = SqlFirewall::new();
230        assert!(fw.check("DELETE FROM users").is_err());
231    }
232
233    #[test]
234    fn test_allows_delete_with_where() {
235        let fw = SqlFirewall::new();
236        assert!(fw.check("DELETE FROM users WHERE id = 1").is_ok());
237    }
238
239    #[test]
240    fn test_blocks_update_without_where() {
241        let fw = SqlFirewall::new();
242        assert!(fw.check("UPDATE users SET name = 'a'").is_err());
243    }
244
245    #[test]
246    fn test_allows_update_with_where() {
247        let fw = SqlFirewall::new();
248        assert!(fw.check("UPDATE users SET name = 'a' WHERE id = 1").is_ok());
249    }
250
251    #[test]
252    fn test_blocks_grant() {
253        let fw = SqlFirewall::new();
254        assert!(fw.check("GRANT SELECT ON users TO app").is_err());
255    }
256
257    #[test]
258    fn test_blocks_revoke() {
259        let fw = SqlFirewall::new();
260        assert!(fw.check("REVOKE SELECT ON users FROM app").is_err());
261    }
262
263    #[test]
264    fn test_blocked_count_increments() {
265        let fw = SqlFirewall::new();
266        assert_eq!(fw.blocked_count(), 0);
267        let _ = fw.check("DROP TABLE x");
268        let _ = fw.check("TRUNCATE TABLE y");
269        assert_eq!(fw.blocked_count(), 2);
270    }
271
272    #[test]
273    fn test_add_custom_rule() {
274        let fw = SqlFirewall::new();
275        fw.add_rule(FirewallRule {
276            name: "block_select_star".into(),
277            pattern: r"(?i)SELECT\s+\*".into(),
278            action: FirewallAction::Block,
279            unless_pattern: None,
280        });
281        // 自定义规则应拦截 SELECT *
282        assert!(fw.check("SELECT * FROM users").is_err());
283        // 普通 SELECT 仍允许
284        assert!(fw.check("SELECT id FROM users").is_ok());
285    }
286
287    #[test]
288    fn test_violation_display() {
289        let v = FirewallViolation {
290            rule_name: "block_drop_table".to_string(),
291            sql: "DROP TABLE users".to_string(),
292            action: FirewallAction::Block,
293        };
294        let s = format!("{}", v);
295        assert!(s.contains("block_drop_table"));
296        assert!(s.contains("DROP TABLE users"));
297    }
298
299    #[test]
300    fn test_case_insensitive_match() {
301        let fw = SqlFirewall::new();
302        assert!(fw.check("drop table users").is_err());
303        assert!(fw.check("Drop Table users").is_err());
304    }
305
306    #[test]
307    fn test_default_impl() {
308        let fw = SqlFirewall::default();
309        assert!(fw.check("SELECT 1").is_ok());
310    }
311}