From 99e5e6317d161f3c71404ee3dd6145c4d855563e Mon Sep 17 00:00:00 2001 From: Illuminated2020 <2357303264@qq.com> Date: Sun, 10 May 2026 17:07:35 +0800 Subject: [PATCH] =?UTF-8?q?fix(sql-safety):=20=E4=BF=AE=E6=AD=A3=E5=8D=B1?= =?UTF-8?q?=E9=99=A9=20SQL=20=E5=88=86=E7=B1=BB=E4=B8=8E=E9=A3=8E=E9=99=A9?= =?UTF-8?q?=E4=B8=8A=E4=B8=8B=E6=96=87?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 扫描可执行 SQL token,避免 WITH、EXPLAIN ANALYZE 和多语句中的写入/DDL 被判为 Read - 引入 RiskContext 支持 environment_label,并保留 risk_for_connection 便捷调用 - 补充 SQL 安全分类与生产环境风险识别测试 BREAKING CHANGE: risk_for 现在接收 RiskContext;三参数便捷调用迁移为 risk_for_connection --- crates/dbx-core/src/sql_safety.rs | 170 ++++++++++++++++++++++++++---- 1 file changed, 152 insertions(+), 18 deletions(-) diff --git a/crates/dbx-core/src/sql_safety.rs b/crates/dbx-core/src/sql_safety.rs index 5f24ada2d..aedf92e3a 100644 --- a/crates/dbx-core/src/sql_safety.rs +++ b/crates/dbx-core/src/sql_safety.rs @@ -28,19 +28,47 @@ pub struct RiskMetadata { pub first_token: Option, } +#[derive(Debug, Clone, Copy)] +pub struct RiskContext<'a> { + pub connection_name: &'a str, + pub color: Option<&'a str>, + pub environment_label: Option<&'a str>, +} + +impl<'a> RiskContext<'a> { + pub fn new(connection_name: &'a str) -> Self { + Self { connection_name, color: None, environment_label: None } + } + + pub fn with_color(mut self, color: Option<&'a str>) -> Self { + self.color = color; + self + } + + pub fn with_environment_label(mut self, environment_label: Option<&'a str>) -> Self { + self.environment_label = environment_label; + self + } +} + pub fn classify_sql(sql: &str) -> OperationClass { - let token = first_executable_token(sql).map(|s| s.to_ascii_uppercase()); - match token.as_deref() { + let tokens = executable_tokens(sql); + if tokens.iter().any(|token| is_ddl_token(token)) { + return OperationClass::Ddl; + } + if tokens.iter().any(|token| is_write_token(token)) { + return OperationClass::Write; + } + + match tokens.first().map(String::as_str) { Some("SELECT" | "SHOW" | "DESCRIBE" | "EXPLAIN" | "WITH") => OperationClass::Read, - Some("INSERT" | "UPDATE" | "DELETE" | "MERGE" | "REPLACE") => OperationClass::Write, - Some("CREATE" | "ALTER" | "DROP" | "TRUNCATE" | "RENAME") => OperationClass::Ddl, _ => OperationClass::Unknown, } } -pub fn risk_for(sql: &str, connection_name: &str, color: Option<&str>) -> RiskMetadata { +pub fn risk_for(sql: &str, context: RiskContext<'_>) -> RiskMetadata { let operation_class = classify_sql(sql); - let (is_production, production_reason) = production_signal(connection_name, color); + let (is_production, production_reason) = production_signal(context); let risk_level = match (operation_class, is_production) { (OperationClass::Read, _) => RiskLevel::Low, (OperationClass::Write, false) => RiskLevel::Medium, @@ -58,22 +86,100 @@ pub fn risk_for(sql: &str, connection_name: &str, color: Option<&str>) -> RiskMe } } -fn production_signal(connection_name: &str, color: Option<&str>) -> (bool, Option) { - if matches!(color, Some("#ef4444")) { +pub fn risk_for_connection(sql: &str, connection_name: &str, color: Option<&str>) -> RiskMetadata { + risk_for(sql, RiskContext::new(connection_name).with_color(color)) +} + +fn production_signal(context: RiskContext<'_>) -> (bool, Option) { + if matches!(context.color, Some("#ef4444")) { return (true, Some("red connection color".to_string())); } - let name = connection_name.to_ascii_lowercase(); - if ["prod", "production", "live"] - .iter() - .any(|needle| name.contains(needle)) - { + if context.environment_label.is_some_and(contains_production_signal) { + return (true, Some("environment label".to_string())); + } + + if contains_production_signal(context.connection_name) { return (true, Some("connection name fallback".to_string())); } (false, None) } +fn contains_production_signal(value: &str) -> bool { + let value = value.to_ascii_lowercase(); + ["prod", "production", "live"].iter().any(|needle| value.contains(needle)) +} + +fn is_write_token(token: &str) -> bool { + matches!(token, "INSERT" | "UPDATE" | "DELETE" | "MERGE" | "REPLACE") +} + +fn is_ddl_token(token: &str) -> bool { + matches!(token, "CREATE" | "ALTER" | "DROP" | "TRUNCATE" | "RENAME") +} + +fn executable_tokens(sql: &str) -> 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';' { + i += 1; + continue; + } + + 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; + } + + if matches!(bytes[i], b'\'' | b'"' | b'`') { + let quote = bytes[i]; + i += 1; + while i < bytes.len() { + if bytes[i] == quote { + if i + 1 < bytes.len() && bytes[i + 1] == quote { + i += 2; + continue; + } + i += 1; + break; + } + i += 1; + } + continue; + } + + if bytes[i].is_ascii_alphabetic() || bytes[i] == b'_' { + let start = i; + i += 1; + while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') { + i += 1; + } + tokens.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; @@ -123,17 +229,45 @@ mod tests { #[test] fn classifies_write_and_ddl() { - assert_eq!( - classify_sql("update users set name = 'a'"), - OperationClass::Write - ); + assert_eq!(classify_sql("update users set name = 'a'"), OperationClass::Write); assert_eq!(classify_sql("DROP TABLE users"), OperationClass::Ddl); } + #[test] + fn with_does_not_hide_write_or_ddl() { + assert_eq!( + classify_sql("WITH moved AS (DELETE FROM orders RETURNING *) SELECT * FROM moved"), + OperationClass::Write + ); + assert_eq!(classify_sql("WITH dropped AS (DROP TABLE old_orders) SELECT 1"), OperationClass::Ddl); + } + + #[test] + fn explain_analyze_write_is_write() { + assert_eq!(classify_sql("EXPLAIN ANALYZE UPDATE users SET name = 'a'"), OperationClass::Write); + } + + #[test] + fn dangerous_statement_in_multi_statement_sql_is_not_read() { + assert_eq!(classify_sql("SELECT * FROM users; DELETE FROM users WHERE id = 1"), OperationClass::Write); + assert_eq!(classify_sql("SHOW TABLES; DROP TABLE users"), OperationClass::Ddl); + } + #[test] fn red_color_marks_production() { - let risk = risk_for("SELECT * FROM orders", "prod-main", Some("#ef4444")); + let risk = risk_for_connection("SELECT * FROM orders", "prod-main", Some("#ef4444")); assert!(risk.is_production); assert_eq!(risk.risk_level, RiskLevel::Low); } + + #[test] + fn environment_label_marks_production() { + let risk = risk_for( + "UPDATE orders SET status = 'done'", + 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); + } }