diff --git a/src/rpc/handlers/query.rs b/src/rpc/handlers/query.rs index 9d11c0f..35a0373 100644 --- a/src/rpc/handlers/query.rs +++ b/src/rpc/handlers/query.rs @@ -72,52 +72,155 @@ pub fn validate_query_or_mutation_id(id: &str, field_name: &str) -> Result<(), D Ok(()) } -fn strip_sql_comments_and_trim(sql: &str) -> String { - let mut res = String::new(); - let mut chars = sql.chars().peekable(); - let mut in_single_comment = false; - let mut in_multi_comment = false; - let mut in_string = false; - let mut string_quote = ' '; +/// Zero-allocation, lazy SQL token scanner used by the Safe Mode precheck and +/// query classification. +/// +/// Yields whitespace-separated tokens as slices of the original text, skipping +/// `-- line` and `/* block */` comments. Quoted literals (`'...'`, `"..."`, +/// with `\\` and doubled-quote escapes) are kept inside a single token, so their +/// contents can never be mistaken for keywords or comment markers. Scanning +/// stops as soon as the caller stops pulling tokens. +struct SqlTokens<'a> { + src: &'a str, + pos: usize, +} - while let Some(c) = chars.next() { - if in_single_comment { - if c == '\n' { - in_single_comment = false; - res.push(' '); +impl<'a> SqlTokens<'a> { + fn new(src: &'a str) -> Self { + Self { src, pos: 0 } + } + + fn char_at(&self, i: usize) -> char { + self.src[i..].chars().next().unwrap_or('\0') + } +} + +impl<'a> Iterator for SqlTokens<'a> { + type Item = &'a str; + + fn next(&mut self) -> Option<&'a str> { + let b = self.src.as_bytes(); + let n = b.len(); + let mut i = self.pos; + + // Skip whitespace and comments before the token. + loop { + if i >= n { + self.pos = n; + return None; } - } else if in_multi_comment { - if c == '*' && chars.peek() == Some(&'/') { - chars.next(); - in_multi_comment = false; - res.push(' '); + if b[i] == b'-' && b.get(i + 1) == Some(&b'-') { + i = self.src[i..].find('\n').map_or(n, |off| i + off); + } else if b[i] == b'/' && b.get(i + 1) == Some(&b'*') { + i = self.src[i + 2..] + .find("*/") + .map_or(n, |off| i + 2 + off + 2); + } else { + let c = self.char_at(i); + if !c.is_whitespace() { + break; + } + i += c.len_utf8(); } - } else if in_string { - if c == '\\' { - chars.next(); - } else if c == string_quote { - if chars.peek() == Some(&string_quote) { - chars.next(); - } else { - in_string = false; - res.push('\''); + } + + let start = i; + while i < n { + match b[i] { + b'-' if b.get(i + 1) == Some(&b'-') => break, + b'/' if b.get(i + 1) == Some(&b'*') => break, + q @ (b'\'' | b'"') => { + i += 1; + while i < n { + if b[i] == b'\\' { + i += 1; + if i < n { + i += self.char_at(i).len_utf8(); + } + } else if b[i] == q { + i += 1; + if b.get(i) == Some(&q) { + i += 1; + } else { + break; + } + } else { + i += 1; + } + } + } + _ => { + let c = self.char_at(i); + if c.is_whitespace() { + break; + } + i += c.len_utf8(); } } - } else if c == '-' && chars.peek() == Some(&'-') { - chars.next(); - in_single_comment = true; - } else if c == '/' && chars.peek() == Some(&'*') { - chars.next(); - in_multi_comment = true; - } else if c == '\'' || c == '"' { - in_string = true; - string_quote = c; - res.push('\''); - } else { - res.push(c); } + + let i = i.min(n); + self.pos = i; + Some(&self.src[start..i]) } - res.trim().to_uppercase() +} + +/// Case-insensitive (ASCII) check that the first SQL token starts with `prefix`. +fn first_token_starts_with(sql: &str, prefix: &str) -> bool { + SqlTokens::new(sql).next().is_some_and(|t| { + t.get(..prefix.len()) + .is_some_and(|head| head.eq_ignore_ascii_case(prefix)) + }) +} + +const KW_DROP: u16 = 1 << 0; +const KW_TRUNCATE: u16 = 1 << 1; +const KW_DELETE: u16 = 1 << 2; +const KW_UPDATE: u16 = 1 << 3; +const KW_ALTER: u16 = 1 << 4; +const KW_TABLE: u16 = 1 << 5; +const KW_DATABASE: u16 = 1 << 6; +const KW_INSERT: u16 = 1 << 7; +const KW_INTO: u16 = 1 << 8; +const KW_CREATE: u16 = 1 << 9; +const KW_VALUES: u16 = 1 << 10; +const KW_SELECT: u16 = 1 << 11; +/// A mutating `ALTER TABLE` action keyword (parentheses ignored). +const KW_ALTER_ACTION: u16 = 1 << 12; + +/// Single pass over all tokens, recording which dangerous keywords appear. +fn scan_keywords(sql: &str) -> u16 { + const EXACT: [(&str, u16); 12] = [ + ("DROP", KW_DROP), + ("TRUNCATE", KW_TRUNCATE), + ("DELETE", KW_DELETE), + ("UPDATE", KW_UPDATE), + ("ALTER", KW_ALTER), + ("TABLE", KW_TABLE), + ("DATABASE", KW_DATABASE), + ("INSERT", KW_INSERT), + ("INTO", KW_INTO), + ("CREATE", KW_CREATE), + ("VALUES", KW_VALUES), + ("SELECT", KW_SELECT), + ]; + const ALTER_ACTIONS: [&str; 9] = [ + "DROP", "DELETE", "UPDATE", "MODIFY", "REPLACE", "CLEAR", "FREEZE", "ATTACH", "DETACH", + ]; + + let mut flags = 0u16; + for token in SqlTokens::new(sql) { + for (kw, bit) in EXACT { + if token.eq_ignore_ascii_case(kw) { + flags |= bit; + } + } + let clean = token.trim_matches(|c| c == '(' || c == ')'); + if ALTER_ACTIONS.iter().any(|a| clean.eq_ignore_ascii_case(a)) { + flags |= KW_ALTER_ACTION; + } + } + flags } /// Splits SQL text into individual statements separated by semicolon (`;`), @@ -189,70 +292,68 @@ pub fn split_sql_statements(sql: &str) -> Vec { /// Checks an individual SQL statement for dangerous/destructive operations in Safe Mode. fn check_single_statement_for_safe_mode(statement_sql: &str) -> Result<(), DriverError> { - let upper = strip_sql_comments_and_trim(statement_sql); - let tokens: Vec<&str> = upper.split_whitespace().collect(); - if tokens.is_empty() { + let mut tokens = SqlTokens::new(statement_sql); + let Some(first) = tokens.next() else { + return Ok(()); + }; + let first = first.trim_start_matches('('); + let is_first = |kw: &str| first.eq_ignore_ascii_case(kw); + + // Fast path: read-only leading commands need no further scanning. + const READ_ONLY: [&str; 8] = [ + "SELECT", "SHOW", "DESCRIBE", "DESC", "EXPLAIN", "EXISTS", "CHECK", "WITH", + ]; + if READ_ONLY.iter().any(|kw| is_first(kw)) { return Ok(()); } + if is_first("UPDATE") { + return Err(safe_mode_violation()); + } - let first = tokens[0].trim_start_matches('('); - let second = tokens.get(1).copied().unwrap_or("").trim_start_matches('('); - let third = tokens.get(2).copied().unwrap_or("").trim_start_matches('('); + let second = tokens.next().unwrap_or("").trim_start_matches('('); + let third = tokens.next().unwrap_or("").trim_start_matches('('); + let second_is = |kw: &str| second.eq_ignore_ascii_case(kw); + let third_is = |kw: &str| third.eq_ignore_ascii_case(kw); + let is_object_kind = || { + second_is("DATABASE") || second_is("TABLE") || second_is("VIEW") || second_is("DICTIONARY") + }; - let is_dangerous = match first { - "SELECT" | "SHOW" | "DESCRIBE" | "DESC" | "EXPLAIN" | "EXISTS" | "CHECK" | "WITH" => false, - "DROP" => { - second == "DATABASE" || second == "TABLE" || second == "VIEW" || second == "DICTIONARY" - } - "TRUNCATE" => second == "TABLE", - "ALTER" => { - second == "TABLE" - && tokens.iter().any(|&t| { - let clean = t.trim_matches(|c| c == '(' || c == ')'); - clean == "DROP" - || clean == "DELETE" - || clean == "UPDATE" - || clean == "MODIFY" - || clean == "REPLACE" - || clean == "CLEAR" - || clean == "FREEZE" - || clean == "ATTACH" - || clean == "DETACH" - }) - } - "INSERT" => { - second == "INTO" - || third == "INTO" - || tokens.contains(&"VALUES") - || tokens.contains(&"SELECT") - } - "DELETE" => second == "FROM" || third == "FROM", - "UPDATE" => true, - "CREATE" => { - second == "DATABASE" || second == "TABLE" || second == "VIEW" || second == "DICTIONARY" - } - "RENAME" => second == "TABLE" || second == "DATABASE", - "ATTACH" | "DETACH" => second == "TABLE" || second == "PARTITION", - _ => { - tokens.contains(&"DROP") - || tokens.contains(&"TRUNCATE") - || tokens.contains(&"DELETE") - || tokens.contains(&"UPDATE") - || (tokens.contains(&"ALTER") && tokens.contains(&"TABLE")) - || (tokens.contains(&"INSERT") && tokens.contains(&"INTO")) - || (tokens.contains(&"CREATE") - && (tokens.contains(&"TABLE") || tokens.contains(&"DATABASE"))) - } + let is_dangerous = if is_first("DROP") || is_first("CREATE") { + is_object_kind() + } else if is_first("TRUNCATE") { + second_is("TABLE") + } else if is_first("ALTER") { + second_is("TABLE") && scan_keywords(statement_sql) & KW_ALTER_ACTION != 0 + } else if is_first("INSERT") { + second_is("INTO") + || third_is("INTO") + || scan_keywords(statement_sql) & (KW_VALUES | KW_SELECT) != 0 + } else if is_first("DELETE") { + second_is("FROM") || third_is("FROM") + } else if is_first("RENAME") { + second_is("TABLE") || second_is("DATABASE") + } else if is_first("ATTACH") || is_first("DETACH") { + second_is("TABLE") || second_is("PARTITION") + } else { + let f = scan_keywords(statement_sql); + f & (KW_DROP | KW_TRUNCATE | KW_DELETE | KW_UPDATE) != 0 + || (f & KW_ALTER != 0 && f & KW_TABLE != 0) + || (f & KW_INSERT != 0 && f & KW_INTO != 0) + || (f & KW_CREATE != 0 && f & (KW_TABLE | KW_DATABASE) != 0) }; if is_dangerous { - return Err(DriverError::SafeModeViolation( - "Operation blocked by Safe Mode: write or destructive queries are forbidden in analytical read-only mode".to_string(), - )); + return Err(safe_mode_violation()); } Ok(()) } +fn safe_mode_violation() -> DriverError { + DriverError::SafeModeViolation( + "Operation blocked by Safe Mode: write or destructive queries are forbidden in analytical read-only mode".to_string(), + ) +} + /// Pre-checks AST/SQL syntax in Safe Mode (`readonly = true`) before network roundtrip, /// evaluating all statements in multi-statement queries. fn enforce_safe_mode_precheck(sql: &str) -> Result<(), DriverError> { @@ -297,14 +398,12 @@ pub async fn handle_query(params: Option) -> Result { // Classify on the comment-stripped statement so a leading `-- comment` or // `/* comment */` doesn't hide the real starting keyword, and recognize // `WITH ...` CTE queries as tabular too. - let normalized_sql = strip_sql_comments_and_trim(trimmed_sql); - let is_tabular_query = normalized_sql.starts_with("SELECT") - || normalized_sql.starts_with("SHOW") - || normalized_sql.starts_with("DESCRIBE") - || normalized_sql.starts_with("EXPLAIN") - || normalized_sql.starts_with("WITH"); - - let sql_to_run = if is_tabular_query && !normalized_sql.contains("FORMAT ") { + let is_tabular_query = ["SELECT", "SHOW", "DESCRIBE", "EXPLAIN", "WITH"] + .iter() + .any(|kw| first_token_starts_with(trimmed_sql, kw)); + let has_format_clause = SqlTokens::new(trimmed_sql).any(|t| t.eq_ignore_ascii_case("FORMAT")); + + let sql_to_run = if is_tabular_query && !has_format_clause { // FORMAT must precede the statement-terminating `;` in ClickHouse's // grammar, so strip any trailing semicolon before appending it. let sql_no_trailing_semicolon = @@ -683,21 +782,41 @@ mod tests { } #[test] - fn test_strip_sql_comments_handles_escaped_and_doubled_quotes() { + fn test_sql_tokens_handles_escaped_and_doubled_quotes() { // Regression for issue #59: an escaped quote (`\'`) must not prematurely // close a string literal and expose a following `--` as a real comment. - let escaped = - strip_sql_comments_and_trim("SELECT 'Customer\\'s notes -- internal' FROM feedback"); - assert!(escaped.starts_with("SELECT")); - assert!(escaped.ends_with("FROM FEEDBACK")); + let escaped: Vec<&str> = + SqlTokens::new("SELECT 'Customer\\'s notes -- internal' FROM feedback").collect(); + assert_eq!(escaped.first(), Some(&"SELECT")); + assert_eq!(escaped.last(), Some(&"feedback")); + assert!(escaped.contains(&"FROM")); // A doubled quote (`''`, the SQL-standard escape) must not close the - // string either, so the literal's content never leaks into the - // normalized output as bare keywords. - let doubled = strip_sql_comments_and_trim("SELECT 'Don''t drop table' FROM logs"); - assert!(doubled.starts_with("SELECT")); - assert!(doubled.ends_with("FROM LOGS")); - assert!(!doubled.contains("DROP TABLE")); + // string either, so the literal's content never leaks out as keywords. + let doubled: Vec<&str> = SqlTokens::new("SELECT 'Don''t drop table' FROM logs").collect(); + assert_eq!(doubled.first(), Some(&"SELECT")); + assert_eq!(doubled.last(), Some(&"logs")); + assert!(!doubled.iter().any(|t| t.eq_ignore_ascii_case("DROP"))); + } + + #[test] + fn test_safe_mode_precheck_is_case_insensitive_and_comment_aware() { + assert!(enforce_safe_mode_precheck("-- hi\n/* x */ drop table t").is_err()); + assert!(enforce_safe_mode_precheck("Alter Table t Delete where 1").is_err()); + assert!(enforce_safe_mode_precheck("insert into t values (1)").is_err()); + assert!(enforce_safe_mode_precheck("select 'drop table t' -- drop table x").is_ok()); + assert!(enforce_safe_mode_precheck(" \n ").is_ok()); + } + + #[test] + fn test_sql_tokens_skips_comments_and_whitespace() { + let tokens: Vec<&str> = + SqlTokens::new(" -- lead\n/* block */ WITH/**/x AS (SELECT 1) -- tail").collect(); + assert_eq!(tokens, ["WITH", "x", "AS", "(SELECT", "1)"]); + assert!(SqlTokens::new("-- only comment").next().is_none()); + assert!(SqlTokens::new("/* unterminated").next().is_none()); + assert!(first_token_starts_with(" /* c */ select 1", "SELECT")); + assert!(!first_token_starts_with("(select 1)", "SELECT")); } #[test] @@ -846,8 +965,7 @@ mod tests { fn test_tabular_query_trailing_semicolon_format_placement() { // The FORMAT clause must be appended before any trailing `;`, never after. let trimmed_sql = "SELECT 1;"; - let normalized_sql = strip_sql_comments_and_trim(trimmed_sql); - assert!(normalized_sql.starts_with("SELECT")); + assert!(first_token_starts_with(trimmed_sql, "SELECT")); let sql_no_trailing_semicolon = trimmed_sql.trim_end_matches(|c: char| c == ';' || c.is_whitespace()); let sql_to_run = format!(