1use std::sync::RwLock;
4
5#[derive(Debug, Clone)]
7pub struct FirewallRule {
8 pub name: String,
10 pub pattern: String,
12 pub action: FirewallAction,
14 pub unless_pattern: Option<String>,
16}
17
18#[derive(Debug, Clone, PartialEq, Eq)]
20pub enum FirewallAction {
21 Block,
23 Log,
25 RequireApproval,
27}
28
29pub 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 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 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 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 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 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 pub fn blocked_count(&self) -> u64 {
151 self.blocked_count
152 .load(std::sync::atomic::Ordering::Relaxed)
153 }
154
155 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#[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 assert!(fw.check("SELECT * FROM users").is_err());
283 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}