From 99e79976209b6aee3434746dcb5ed8a68a9db594 Mon Sep 17 00:00:00 2001 From: Illuminated2020 <2357303264@qq.com> Date: Sun, 10 May 2026 17:19:15 +0800 Subject: [PATCH] =?UTF-8?q?fix(sql=5Fsafety):=20=E4=BF=AE=E6=AD=A3?= =?UTF-8?q?=E5=AD=90=E6=9F=A5=E8=AF=A2=E8=BE=B9=E7=95=8C=E8=AF=AF=E5=88=A4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/dbx-core/src/sql_safety.rs | 81 +++++++++++++++++++++++++++---- 1 file changed, 72 insertions(+), 9 deletions(-) diff --git a/crates/dbx-core/src/sql_safety.rs b/crates/dbx-core/src/sql_safety.rs index c8fb3d634..bc56db6d4 100644 --- a/crates/dbx-core/src/sql_safety.rs +++ b/crates/dbx-core/src/sql_safety.rs @@ -150,14 +150,22 @@ fn is_ddl_token(token: &str) -> bool { matches!(token, "CREATE" | "ALTER" | "DROP" | "TRUNCATE" | "RENAME") } +#[derive(Debug, Clone, PartialEq, Eq)] +struct SqlToken { + text: String, + depth: usize, +} + fn has_unfiltered_destructive_write(sql: &str) -> bool { - executable_statements(sql).into_iter().any(|statement| { + scanned_executable_statements(sql).into_iter().any(|statement| { statement.iter().enumerate().any(|(index, token)| { - if !matches!(token.as_str(), "DELETE" | "UPDATE") { + if !matches!(token.text.as_str(), "DELETE" | "UPDATE") { return false; } - !statement[index + 1..].iter().any(|token| matches!(token.as_str(), "WHERE" | "LIMIT")) + !statement[index + 1..] + .iter() + .any(|boundary| boundary.depth == token.depth && matches!(boundary.text.as_str(), "WHERE" | "LIMIT")) }) }) } @@ -167,14 +175,27 @@ fn executable_tokens(sql: &str) -> Vec { } fn executable_statements(sql: &str) -> Vec> { + scanned_executable_statements(sql) + .into_iter() + .map(|statement| statement.into_iter().map(|token| token.text).collect()) + .collect() +} + +fn scanned_executable_statements(sql: &str) -> Vec> { let mut statements = Vec::new(); let mut current = Vec::new(); - scan_executable_tokens(sql, &mut current, &mut statements); + let mut depth = 0; + scan_executable_tokens(sql, &mut current, &mut statements, &mut depth); push_statement(&mut current, &mut statements); statements } -fn scan_executable_tokens(sql: &str, current: &mut Vec, statements: &mut Vec>) { +fn scan_executable_tokens( + sql: &str, + current: &mut Vec, + statements: &mut Vec>, + depth: &mut usize, +) { let bytes = sql.as_bytes(); let mut i = 0; @@ -185,7 +206,21 @@ fn scan_executable_tokens(sql: &str, current: &mut Vec, statements: &mut } if bytes[i] == b';' { - push_statement(current, statements); + if *depth == 0 { + push_statement(current, statements); + } + i += 1; + continue; + } + + if bytes[i] == b'(' { + *depth += 1; + i += 1; + continue; + } + + if bytes[i] == b')' { + *depth = depth.saturating_sub(1); i += 1; continue; } @@ -202,7 +237,7 @@ fn scan_executable_tokens(sql: &str, current: &mut Vec, statements: &mut if i + 2 < bytes.len() && bytes[i + 2] == b'!' { let content_start = i + 3; let content_end = block_comment_end(bytes, content_start); - scan_executable_tokens(&sql[content_start..content_end], current, statements); + scan_executable_tokens(&sql[content_start..content_end], current, statements, depth); i = (content_end + 2).min(bytes.len()); } else { i += 2; @@ -248,7 +283,7 @@ fn scan_executable_tokens(sql: &str, current: &mut Vec, statements: &mut while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') { i += 1; } - current.push(sql[start..i].to_ascii_uppercase()); + current.push(SqlToken { text: sql[start..i].to_ascii_uppercase(), depth: *depth }); continue; } @@ -256,7 +291,7 @@ fn scan_executable_tokens(sql: &str, current: &mut Vec, statements: &mut } } -fn push_statement(current: &mut Vec, statements: &mut Vec>) { +fn push_statement(current: &mut Vec, statements: &mut Vec>) { if !current.is_empty() { statements.push(std::mem::take(current)); } @@ -380,6 +415,34 @@ mod tests { assert_eq!(risk_for("TRUNCATE TABLE users", RiskContext::new("dev")).risk_level, RiskLevel::Critical); } + #[test] + fn destructive_writes_only_count_top_level_where_or_limit_as_boundaries() { + assert_eq!( + risk_for( + "DELETE FROM users USING (SELECT id FROM archived WHERE stale = true) old", + RiskContext::new("dev") + ) + .risk_level, + RiskLevel::Critical + ); + assert_eq!( + risk_for( + "UPDATE users SET active = false FROM (SELECT id FROM flags LIMIT 10) flags", + RiskContext::new("dev") + ) + .risk_level, + RiskLevel::Critical + ); + assert_eq!( + risk_for( + "DELETE FROM users WHERE id IN (SELECT user_id FROM archived WHERE stale = true)", + RiskContext::new("dev") + ) + .risk_level, + RiskLevel::Medium + ); + } + #[test] fn postgresql_dollar_quotes_do_not_contribute_tokens() { assert_eq!(classify_sql("SELECT $$ DELETE FROM users $$"), OperationClass::Read);