diff --git a/crates/dbx-core/src/sql_safety.rs b/crates/dbx-core/src/sql_safety.rs index aedf92e3a..c8fb3d634 100644 --- a/crates/dbx-core/src/sql_safety.rs +++ b/crates/dbx-core/src/sql_safety.rs @@ -53,6 +53,10 @@ impl<'a> RiskContext<'a> { pub fn classify_sql(sql: &str) -> OperationClass { let tokens = executable_tokens(sql); + classify_tokens(&tokens) +} + +fn classify_tokens(tokens: &[String]) -> OperationClass { if tokens.iter().any(|token| is_ddl_token(token)) { return OperationClass::Ddl; } @@ -71,6 +75,7 @@ pub fn risk_for(sql: &str, context: RiskContext<'_>) -> RiskMetadata { let (is_production, production_reason) = production_signal(context); let risk_level = match (operation_class, is_production) { (OperationClass::Read, _) => RiskLevel::Low, + (OperationClass::Write, _) if has_unfiltered_destructive_write(sql) => RiskLevel::Critical, (OperationClass::Write, false) => RiskLevel::Medium, (OperationClass::Write, true) => RiskLevel::High, (OperationClass::Ddl, _) => RiskLevel::Critical, @@ -82,7 +87,7 @@ pub fn risk_for(sql: &str, context: RiskContext<'_>) -> RiskMetadata { risk_level, is_production, production_reason, - first_token: first_executable_token(sql).map(str::to_string), + first_token: first_executable_token(sql), } } @@ -91,12 +96,17 @@ pub fn risk_for_connection(sql: &str, connection_name: &str, color: Option<&str> } fn production_signal(context: RiskContext<'_>) -> (bool, Option) { - if matches!(context.color, Some("#ef4444")) { - return (true, Some("red connection color".to_string())); + if let Some(environment_label) = context.environment_label { + if contains_non_production_signal(environment_label) { + return (false, None); + } + if contains_production_signal(environment_label) { + return (true, Some("environment label".to_string())); + } } - if context.environment_label.is_some_and(contains_production_signal) { - return (true, Some("environment label".to_string())); + if matches!(context.color, Some("#ef4444")) { + return (true, Some("red connection color".to_string())); } if contains_production_signal(context.connection_name) { @@ -111,6 +121,27 @@ fn contains_production_signal(value: &str) -> bool { ["prod", "production", "live"].iter().any(|needle| value.contains(needle)) } +fn contains_non_production_signal(value: &str) -> bool { + let value = value.to_ascii_lowercase(); + [ + "dev", + "development", + "test", + "testing", + "qa", + "stage", + "staging", + "local", + "sandbox", + "non-prod", + "non-production", + "non production", + "nonprod", + ] + .iter() + .any(|needle| value.contains(needle)) +} + fn is_write_token(token: &str) -> bool { matches!(token, "INSERT" | "UPDATE" | "DELETE" | "MERGE" | "REPLACE") } @@ -119,13 +150,42 @@ fn is_ddl_token(token: &str) -> bool { matches!(token, "CREATE" | "ALTER" | "DROP" | "TRUNCATE" | "RENAME") } +fn has_unfiltered_destructive_write(sql: &str) -> bool { + executable_statements(sql).into_iter().any(|statement| { + statement.iter().enumerate().any(|(index, token)| { + if !matches!(token.as_str(), "DELETE" | "UPDATE") { + return false; + } + + !statement[index + 1..].iter().any(|token| matches!(token.as_str(), "WHERE" | "LIMIT")) + }) + }) +} + fn executable_tokens(sql: &str) -> Vec { + executable_statements(sql).into_iter().flatten().collect() +} + +fn executable_statements(sql: &str) -> Vec> { + let mut statements = Vec::new(); + let mut current = Vec::new(); + scan_executable_tokens(sql, &mut current, &mut statements); + push_statement(&mut current, &mut statements); + statements +} + +fn scan_executable_tokens(sql: &str, current: &mut Vec, statements: &mut Vec>) { let bytes = sql.as_bytes(); let mut i = 0; - let mut tokens = Vec::new(); while i < bytes.len() { - if bytes[i].is_ascii_whitespace() || bytes[i] == b';' { + if bytes[i].is_ascii_whitespace() { + i += 1; + continue; + } + + if bytes[i] == b';' { + push_statement(current, statements); i += 1; continue; } @@ -139,11 +199,29 @@ fn executable_tokens(sql: &str) -> Vec { } if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' { - i += 2; - while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') { - i += 1; + 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); + i = (content_end + 2).min(bytes.len()); + } else { + i += 2; + while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') { + i += 1; + } + i = (i + 2).min(bytes.len()); + } + continue; + } + + if let Some(delimiter_len) = dollar_quote_delimiter_len(bytes, i) { + let delimiter = &sql[i..i + delimiter_len]; + i += delimiter_len; + if let Some(end) = sql[i..].find(delimiter) { + i += end + delimiter_len; + } else { + i = bytes.len(); } - i = (i + 2).min(bytes.len()); continue; } @@ -170,51 +248,54 @@ fn executable_tokens(sql: &str) -> Vec { while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') { i += 1; } - tokens.push(sql[start..i].to_ascii_uppercase()); + current.push(sql[start..i].to_ascii_uppercase()); continue; } i += 1; } - - tokens } -fn first_executable_token(sql: &str) -> Option<&str> { - let bytes = sql.as_bytes(); - let mut i = 0; +fn push_statement(current: &mut Vec, statements: &mut Vec>) { + if !current.is_empty() { + statements.push(std::mem::take(current)); + } +} - while i < bytes.len() { - while i < bytes.len() && bytes[i].is_ascii_whitespace() { - i += 1; +fn block_comment_end(bytes: &[u8], mut i: usize) -> usize { + while i + 1 < bytes.len() { + if bytes[i] == b'*' && bytes[i + 1] == b'/' { + return i; } + i += 1; + } + bytes.len() +} - if i + 1 < bytes.len() && bytes[i] == b'-' && bytes[i + 1] == b'-' { - i += 2; - while i < bytes.len() && bytes[i] != b'\n' { - i += 1; - } - continue; - } - - if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' { - i += 2; - while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') { - i += 1; - } - i = (i + 2).min(bytes.len()); - continue; - } - - break; +fn dollar_quote_delimiter_len(bytes: &[u8], start: usize) -> Option { + if bytes.get(start) != Some(&b'$') { + return None; } - let start = i; - while i < bytes.len() && (bytes[i].is_ascii_alphabetic() || bytes[i] == b'_') { + let mut i = start + 1; + if bytes.get(i) == Some(&b'$') { + return Some(2); + } + + if !bytes.get(i).is_some_and(|byte| byte.is_ascii_alphabetic() || *byte == b'_') { + return None; + } + + i += 1; + while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') { i += 1; } - (i > start).then_some(&sql[start..i]) + (bytes.get(i) == Some(&b'$')).then_some(i - start + 1) +} + +fn first_executable_token(sql: &str) -> Option { + executable_tokens(sql).into_iter().next() } #[cfg(test)] @@ -263,11 +344,53 @@ mod tests { #[test] fn environment_label_marks_production() { let risk = risk_for( - "UPDATE orders SET status = 'done'", + "UPDATE orders SET status = 'done' WHERE id = 1", RiskContext { connection_name: "analytics", color: None, environment_label: Some("Production") }, ); assert!(risk.is_production); assert_eq!(risk.production_reason.as_deref(), Some("environment label")); assert_eq!(risk.risk_level, RiskLevel::High); } + + #[test] + fn environment_label_overrides_color_and_name_fallback() { + let non_prod_label = risk_for( + "SELECT * FROM orders", + RiskContext { connection_name: "prod-main", color: Some("#ef4444"), environment_label: Some("Staging") }, + ); + assert!(!non_prod_label.is_production); + assert_eq!(non_prod_label.production_reason, None); + + let prod_label = risk_for( + "SELECT * FROM orders", + RiskContext { connection_name: "analytics", color: Some("#22c55e"), environment_label: Some("Production") }, + ); + assert!(prod_label.is_production); + assert_eq!(prod_label.production_reason.as_deref(), Some("environment label")); + } + + #[test] + fn destructive_writes_without_where_or_limit_are_critical() { + assert_eq!(risk_for("DELETE FROM users", RiskContext::new("dev")).risk_level, RiskLevel::Critical); + assert_eq!( + risk_for("UPDATE users SET active = false", RiskContext::new("dev")).risk_level, + RiskLevel::Critical + ); + assert_eq!(risk_for("DELETE FROM users WHERE id = 1", RiskContext::new("dev")).risk_level, RiskLevel::Medium); + assert_eq!(risk_for("TRUNCATE TABLE users", RiskContext::new("dev")).risk_level, RiskLevel::Critical); + } + + #[test] + fn postgresql_dollar_quotes_do_not_contribute_tokens() { + assert_eq!(classify_sql("SELECT $$ DELETE FROM users $$"), OperationClass::Read); + assert_eq!(classify_sql("SELECT $tag$ DROP TABLE users $tag$"), OperationClass::Read); + } + + #[test] + fn mysql_executable_comment_contributes_tokens() { + assert_eq!(classify_sql("/*!50000 DELETE FROM users */ SELECT 1"), OperationClass::Write); + let risk = risk_for("/*! UPDATE users SET active = false */", RiskContext::new("dev")); + assert_eq!(risk.operation_class, OperationClass::Write); + assert_eq!(risk.first_token.as_deref(), Some("UPDATE")); + } }