From 1f666f99b3fbfc7b4d955caf51e7af244b27de26 Mon Sep 17 00:00:00 2001 From: t8y2 <1156263951@qq.com> Date: Sun, 24 May 2026 10:26:09 +0800 Subject: [PATCH] fix(duckdb): support multiple schemas beyond default "main" DuckDB was not recognized as a schema-aware database, so only the default "main" schema was shown in the object tree and schema selector. Added DuckDB to schema-aware type lists, implemented schema listing via information_schema.schemata, and replaced hardcoded table_schema = 'main' with parameterized queries. --- .../desktop/src/lib/databaseCapabilitySets.ts | 2 + crates/dbx-core/src/query.rs | 12 +- crates/dbx-core/src/schema.rs | 70 ++++-- crates/dbx-core/src/sql.rs | 219 ++++++++++++++++-- crates/dbx-core/src/sql_dialect.rs | 1 + crates/dbx-web/src/routes/query.rs | 15 +- src-tauri/src/commands/query.rs | 21 +- 7 files changed, 294 insertions(+), 46 deletions(-) diff --git a/apps/desktop/src/lib/databaseCapabilitySets.ts b/apps/desktop/src/lib/databaseCapabilitySets.ts index 1f0c2a21d..45c5b2569 100644 --- a/apps/desktop/src/lib/databaseCapabilitySets.ts +++ b/apps/desktop/src/lib/databaseCapabilitySets.ts @@ -25,6 +25,7 @@ export const SCHEMA_AWARE_TYPES = new Set([ "trino", "db2", "tdengine", + "duckdb", ]); export const SQL_FILE_UNSUPPORTED_TYPES = new Set(["redis", "mongodb", "elasticsearch"]); @@ -201,6 +202,7 @@ export const TREE_SCHEMA_TYPES = new Set([ "trino", "h2", "tdengine", + "duckdb", ]); export const PG_LIKE_STRUCTURE_TYPES = new Set(["postgres", "redshift", "gaussdb", "opengauss"]); diff --git a/crates/dbx-core/src/query.rs b/crates/dbx-core/src/query.rs index 1c0b7c429..ad57b2e14 100644 --- a/crates/dbx-core/src/query.rs +++ b/crates/dbx-core/src/query.rs @@ -8,12 +8,18 @@ use tokio_util::sync::CancellationToken; use crate::connection::{AppState, PoolKind}; use crate::db; +use crate::models::connection::DatabaseType; use crate::sql::{split_sql_batches, split_sql_statements, starts_with_executable_sql_keyword}; pub const QUERY_TIMEOUT: Duration = Duration::from_secs(30); pub const MAX_ROWS: usize = 10000; pub const QUERY_CANCELED: &str = "Query canceled"; +async fn connection_database_type(state: &AppState, connection_id: &str) -> Option { + let configs = state.configs.read().await; + configs.get(connection_id).map(|config| config.db_type) +} + #[derive(Clone, Debug, Default)] pub struct QueryExecutionOptions { pub max_rows: Option, @@ -683,7 +689,11 @@ pub async fn execute_multi_core_with_options( return execute_multi_sqlserver(state, &pool_key, sql, cancel_token, options).await; } - let statements = split_sql_statements(sql); + let db_type = connection_database_type(state, connection_id).await; + let statements = db_type.map_or_else( + || split_sql_statements(sql), + |db_type| crate::sql::split_sql_statements_for_database(sql, db_type), + ); if statements.len() <= 1 { let single_sql = statements.into_iter().next().unwrap_or_default(); let result = execute_sql_statement_with_options( diff --git a/crates/dbx-core/src/schema.rs b/crates/dbx-core/src/schema.rs index f065c006a..8473437ad 100644 --- a/crates/dbx-core/src/schema.rs +++ b/crates/dbx-core/src/schema.rs @@ -6,24 +6,29 @@ use crate::db; use crate::models::connection::DatabaseType; pub fn duckdb_query_tables(con: &duckdb::Connection) -> Result, String> { - duckdb_query_tables_in_database(con, "main") + duckdb_query_tables_in_database(con, "main", "main") } -pub fn duckdb_query_tables_in_database(con: &duckdb::Connection, database: &str) -> Result, String> { - duckdb_query_tables_in_database_with_attached(con, database, &[]) +pub fn duckdb_query_tables_in_database( + con: &duckdb::Connection, + database: &str, + schema: &str, +) -> Result, String> { + duckdb_query_tables_in_database_with_attached(con, database, schema, &[]) } pub fn duckdb_query_tables_in_database_with_attached( con: &duckdb::Connection, database: &str, + schema: &str, attached_names: &[String], ) -> Result, String> { let database = duckdb_catalog_name(con, database, attached_names)?; let mut stmt = con.prepare( - "SELECT table_name, table_type FROM information_schema.tables WHERE table_catalog = ? AND table_schema = 'main' ORDER BY table_name" + "SELECT table_name, table_type FROM information_schema.tables WHERE table_catalog = ? AND table_schema = ? ORDER BY table_name" ).map_err(|e| e.to_string())?; let rows = stmt - .query_map([database.as_str()], |row| { + .query_map((database.as_str(), schema), |row| { Ok(db::TableInfo { name: row.get::<_, String>(0)?, table_type: row.get::<_, String>(1)?, comment: None }) }) .map_err(|e| e.to_string())?; @@ -59,6 +64,25 @@ pub fn duckdb_list_databases_with_attached( Ok(rows.filter_map(|row| row.ok()).collect()) } +pub fn duckdb_list_schemas(con: &duckdb::Connection, database: &str) -> Result, String> { + duckdb_list_schemas_with_attached(con, database, &[]) +} + +pub fn duckdb_list_schemas_with_attached( + con: &duckdb::Connection, + database: &str, + attached_names: &[String], +) -> Result, String> { + let database = duckdb_catalog_name(con, database, attached_names)?; + let mut stmt = con + .prepare( + "SELECT schema_name FROM information_schema.schemata WHERE catalog_name = ? AND schema_name NOT IN ('information_schema', 'pg_catalog') ORDER BY schema_name", + ) + .map_err(|e| e.to_string())?; + let rows = stmt.query_map([database.as_str()], |row| row.get::<_, String>(0)).map_err(|e| e.to_string())?; + Ok(rows.filter_map(|r| r.ok()).collect()) +} + fn duckdb_catalog_name(con: &duckdb::Connection, database: &str, attached_names: &[String]) -> Result { if database.trim().is_empty() || database == "main" { return duckdb_primary_catalog(con, attached_names); @@ -95,20 +119,22 @@ fn duckdb_quote_string(value: &str) -> String { } pub fn duckdb_query_columns(con: &duckdb::Connection, table: &str) -> Result, String> { - duckdb_query_columns_in_database(con, "main", table) + duckdb_query_columns_in_database(con, "main", "main", table) } pub fn duckdb_query_columns_in_database( con: &duckdb::Connection, database: &str, + schema: &str, table: &str, ) -> Result, String> { - duckdb_query_columns_in_database_with_attached(con, database, table, &[]) + duckdb_query_columns_in_database_with_attached(con, database, schema, table, &[]) } pub fn duckdb_query_columns_in_database_with_attached( con: &duckdb::Connection, database: &str, + schema: &str, table: &str, attached_names: &[String], ) -> Result, String> { @@ -123,25 +149,26 @@ pub fn duckdb_query_columns_in_database_with_attached( AND tc.table_name = kcu.table_name WHERE tc.constraint_type = 'PRIMARY KEY' AND tc.table_catalog = ? - AND tc.table_schema = 'main' + AND tc.table_schema = ? AND tc.table_name = ? ORDER BY kcu.ordinal_position", ) .map_err(|e| e.to_string())?; - let pk_rows = - pk_stmt.query_map([database.as_str(), table], |row| row.get::<_, String>(0)).map_err(|e| e.to_string())?; + let pk_rows = pk_stmt + .query_map((database.as_str(), schema, table), |row| row.get::<_, String>(0)) + .map_err(|e| e.to_string())?; let primary_keys: std::collections::HashSet = pk_rows.filter_map(|r| r.ok()).collect(); let mut stmt = con .prepare( "SELECT column_name, data_type, is_nullable, column_default FROM information_schema.columns - WHERE table_catalog = ? AND table_schema = 'main' AND table_name = ? + WHERE table_catalog = ? AND table_schema = ? AND table_name = ? ORDER BY ordinal_position", ) .map_err(|e| e.to_string())?; let rows = stmt - .query_map([database.as_str(), table], |row| { + .query_map((database.as_str(), schema, table), |row| { let name = row.get::<_, String>(0)?; Ok(db::ColumnInfo { is_primary_key: primary_keys.contains(&name), @@ -321,6 +348,11 @@ pub async fn list_schemas_core(state: &AppState, connection_id: &str, database: match pool { PoolKind::Postgres(p) => db::postgres::list_schemas(p).await, + PoolKind::DuckDb(con) => { + let duckdb_attached_names = duckdb_attached_database_names(state, connection_id).await; + let con = con.lock().map_err(|e| e.to_string())?; + duckdb_list_schemas_with_attached(&con, database, &duckdb_attached_names) + } _ => Ok(vec![]), } } @@ -362,7 +394,7 @@ pub async fn list_tables_core( if let Some(con) = extract_duckdb(&connections, &pool_key) { drop(connections); let con = con.lock().map_err(|e| e.to_string())?; - return duckdb_query_tables_in_database_with_attached(&con, database, &duckdb_attached_names); + return duckdb_query_tables_in_database_with_attached(&con, database, schema, &duckdb_attached_names); } if let Some(client) = extract_clickhouse(&connections, &pool_key) { drop(connections); @@ -444,8 +476,8 @@ mod tests { duckdb_attach_database(&con, "analytics", path.to_str().unwrap()).unwrap(); con.execute_batch("CREATE TABLE analytics.attached_table(id INTEGER);").unwrap(); - let main_tables = duckdb_query_tables_in_database(&con, "main").unwrap(); - let attached_tables = duckdb_query_tables_in_database(&con, "analytics").unwrap(); + let main_tables = duckdb_query_tables_in_database(&con, "main", "main").unwrap(); + let attached_tables = duckdb_query_tables_in_database(&con, "analytics", "main").unwrap(); assert!(main_tables.iter().any(|table| table.name == "main_table")); assert!(!main_tables.iter().any(|table| table.name == "attached_table")); @@ -588,7 +620,13 @@ pub async fn get_columns_core( if let Some(con) = extract_duckdb(&connections, &pool_key) { drop(connections); let con = con.lock().map_err(|e| e.to_string())?; - return duckdb_query_columns_in_database_with_attached(&con, database, table, &duckdb_attached_names); + return duckdb_query_columns_in_database_with_attached( + &con, + database, + schema, + table, + &duckdb_attached_names, + ); } if let Some(client) = extract_clickhouse(&connections, &pool_key) { drop(connections); diff --git a/crates/dbx-core/src/sql.rs b/crates/dbx-core/src/sql.rs index db882308e..df39b40da 100644 --- a/crates/dbx-core/src/sql.rs +++ b/crates/dbx-core/src/sql.rs @@ -39,6 +39,26 @@ pub enum SqlFileStatementAction { Skip, } +#[derive(Debug, Clone, Copy, Default)] +pub struct SqlParsingOptions { + pub supports_hash_line_comments: bool, +} + +impl SqlParsingOptions { + pub fn for_database_type(db_type: DatabaseType) -> Self { + Self { + supports_hash_line_comments: matches!( + db_type, + DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks | DatabaseType::Goldendb + ), + } + } + + pub fn mysql_compatible() -> Self { + Self { supports_hash_line_comments: true } + } +} + #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct SqlFileProgress { @@ -64,9 +84,14 @@ pub struct SqlStatementSplitter { dollar_quote_tag: Option, previous: Option, custom_delimiter: Option, + options: SqlParsingOptions, } impl SqlStatementSplitter { + pub fn with_options(options: SqlParsingOptions) -> Self { + Self { options, ..Self::default() } + } + pub fn push_chunk(&mut self, chunk: &str) -> Vec { let mut statements = Vec::new(); let chars = chunk.chars().collect::>(); @@ -137,6 +162,13 @@ impl SqlStatementSplitter { i += 1; continue; } + if self.options.supports_hash_line_comments && ch == '#' { + self.in_line_comment = true; + self.buffer.push(ch); + self.previous = Some(ch); + i += 1; + continue; + } if ch == '/' && next == Some('*') { self.in_block_comment = true; self.buffer.push(ch); @@ -188,9 +220,9 @@ impl SqlStatementSplitter { if let Some(new_delim) = parse_delimiter_command(last_line) { self.custom_delimiter = if new_delim == ";" { None } else { Some(new_delim.to_string()) }; if last_line_start > 0 { - let before = self.buffer[..last_line_start].trim(); - if has_executable_sql(before) { - statements.push(before.to_string()); + let before = &self.buffer[..last_line_start]; + if let Some(statement) = executable_sql_slice(before, self.options) { + statements.push(statement.to_string()); } } self.buffer.clear(); @@ -220,8 +252,8 @@ impl SqlStatementSplitter { let last_line = trimmed.rsplit('\n').next().unwrap_or(trimmed).trim(); if parse_delimiter_command(last_line).is_some() { let before = trimmed.rsplitn(2, '\n').nth(1).unwrap_or("").trim(); - if has_executable_sql(before) { - statements.push(before.to_string()); + if let Some(statement) = executable_sql_slice(before, self.options) { + statements.push(statement.to_string()); } self.buffer.clear(); } else if let Some(ref delim) = self.custom_delimiter { @@ -234,8 +266,7 @@ impl SqlStatementSplitter { } fn push_current_statement(&mut self, statements: &mut Vec) { - let statement = self.buffer.trim(); - if has_executable_sql(statement) { + if let Some(statement) = executable_sql_slice(&self.buffer, self.options) { statements.push(statement.to_string()); } self.buffer.clear(); @@ -250,7 +281,15 @@ impl SqlStatementSplitter { } pub fn split_sql_statements(sql: &str) -> Vec { - let mut splitter = SqlStatementSplitter::default(); + split_sql_statements_with_options(sql, SqlParsingOptions::default()) +} + +pub fn split_sql_statements_for_database(sql: &str, db_type: DatabaseType) -> Vec { + split_sql_statements_with_options(sql, SqlParsingOptions::for_database_type(db_type)) +} + +pub fn split_sql_statements_with_options(sql: &str, options: SqlParsingOptions) -> Vec { + let mut splitter = SqlStatementSplitter::with_options(options); let mut statements = splitter.push_chunk(sql); statements.extend(splitter.finish()); statements @@ -264,7 +303,15 @@ pub struct SqlStatementRange { } pub fn find_statement_at_cursor(sql: &str, cursor_pos: usize) -> String { - let statements = split_sql_statement_ranges(sql); + find_statement_at_cursor_with_options(sql, cursor_pos, SqlParsingOptions::default()) +} + +pub fn find_statement_at_cursor_for_database(sql: &str, cursor_pos: usize, db_type: DatabaseType) -> String { + find_statement_at_cursor_with_options(sql, cursor_pos, SqlParsingOptions::for_database_type(db_type)) +} + +pub fn find_statement_at_cursor_with_options(sql: &str, cursor_pos: usize, options: SqlParsingOptions) -> String { + let statements = split_sql_statement_ranges_with_options(sql, options); let cursor = utf16_offset_to_byte_index(sql, cursor_pos); for (idx, statement) in statements.iter().enumerate() { @@ -298,7 +345,12 @@ fn cursor_has_sql_after_cursor_on_line(sql: &str, cursor: usize) -> bool { sql[cursor..line_end].chars().any(|ch| !ch.is_whitespace()) } +#[allow(dead_code)] fn split_sql_statement_ranges(sql: &str) -> Vec { + split_sql_statement_ranges_with_options(sql, SqlParsingOptions::default()) +} + +fn split_sql_statement_ranges_with_options(sql: &str, options: SqlParsingOptions) -> Vec { let mut ranges = Vec::new(); let mut start = 0; let mut i = 0; @@ -348,6 +400,11 @@ fn split_sql_statement_ranges(sql: &str) -> Vec { i += 2; continue; } + if options.supports_hash_line_comments && ch == '#' { + in_line_comment = true; + i += ch.len_utf8(); + continue; + } if ch == '/' && next == Some('*') { in_block_comment = true; i += 2; @@ -365,8 +422,8 @@ fn split_sql_statement_ranges(sql: &str) -> Vec { let line = sql[line_start..i].trim(); if let Some(new_delimiter) = parse_delimiter_command(line) { let before = sql[start..line_start].trim(); - if has_executable_sql(before) { - push_statement_range(&mut ranges, sql, start, line_start); + if has_executable_sql_with_options(before, options) { + push_statement_range(&mut ranges, sql, start, line_start, options); } custom_delimiter = if new_delimiter == ";" { None } else { Some(new_delimiter.to_string()) }; start = i + ch.len_utf8(); @@ -390,7 +447,7 @@ fn split_sql_statement_ranges(sql: &str) -> Vec { i += ch.len_utf8(); } ';' if !in_single_quote && !in_double_quote && !in_backtick && custom_delimiter.is_none() => { - push_statement_range(&mut ranges, sql, start, i); + push_statement_range(&mut ranges, sql, start, i, options); i += ch.len_utf8(); start = i; } @@ -400,7 +457,7 @@ fn split_sql_statement_ranges(sql: &str) -> Vec { if let Some(delimiter) = &custom_delimiter { if sql[start..i].ends_with(delimiter) { let end = i - delimiter.len(); - push_statement_range(&mut ranges, sql, start, end); + push_statement_range(&mut ranges, sql, start, end, options); start = i; } } @@ -413,19 +470,30 @@ fn split_sql_statement_ranges(sql: &str) -> Vec { let last_line = trimmed.rsplit('\n').next().unwrap_or(trimmed).trim(); if parse_delimiter_command(last_line).is_some() { if let Some(line_start) = sql[start..].rfind('\n').map(|pos| start + pos + 1) { - push_statement_range(&mut ranges, sql, start, line_start); + push_statement_range(&mut ranges, sql, start, line_start, options); } } else { - push_statement_range(&mut ranges, sql, start, sql.len()); + push_statement_range(&mut ranges, sql, start, sql.len(), options); } ranges } -fn push_statement_range(ranges: &mut Vec, sql: &str, start: usize, end: usize) { - let text = sql[start..end].trim(); - if has_executable_sql(text) { - ranges.push(SqlStatementRange { text: text.to_string(), start, end }); +fn push_statement_range( + ranges: &mut Vec, + sql: &str, + start: usize, + end: usize, + options: SqlParsingOptions, +) { + let Some((relative_start, relative_end)) = executable_sql_bounds(&sql[start..end], options) else { + return; + }; + let statement_start = start + relative_start; + let statement_end = start + relative_end; + let text = sql[statement_start..statement_end].to_string(); + if !text.is_empty() { + ranges.push(SqlStatementRange { text, start: statement_start, end: statement_end }); } } @@ -584,7 +652,19 @@ pub fn prepare_sql_file_statement( } pub fn starts_with_executable_sql_keyword(sql: &str, keywords: &[&str]) -> bool { - let Some(token) = first_executable_sql_token(sql) else { + starts_with_executable_sql_keyword_with_options(sql, keywords, SqlParsingOptions::default()) +} + +pub fn starts_with_executable_sql_keyword_for_database(sql: &str, keywords: &[&str], db_type: DatabaseType) -> bool { + starts_with_executable_sql_keyword_with_options(sql, keywords, SqlParsingOptions::for_database_type(db_type)) +} + +pub fn starts_with_executable_sql_keyword_with_options( + sql: &str, + keywords: &[&str], + options: SqlParsingOptions, +) -> bool { + let Some(token) = first_executable_sql_token_with_options(sql, options) else { return false; }; keywords.iter().any(|keyword| token.eq_ignore_ascii_case(keyword)) @@ -634,6 +714,14 @@ fn leading_mysql_executable_comment_start(statement: &str) -> Option { continue; } + if bytes[i] == b'#' { + i += 1; + while i < bytes.len() && bytes[i] != b'\n' { + i += 1; + } + continue; + } + if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' { if i + 2 < bytes.len() && (bytes[i + 2] == b'!' || (i + 3 < bytes.len() && &bytes[i + 2..i + 4] == b"M!")) { return Some(i); @@ -672,7 +760,7 @@ fn is_mysql_lock_table_statement(statement: &str) -> bool { } fn is_mysql_session_restore_statement(statement: &str) -> bool { - let executable = leading_executable_sql(statement); + let executable = leading_executable_sql_with_options(statement, SqlParsingOptions::mysql_compatible()); let upper = executable.split_whitespace().collect::>().join(" ").to_ascii_uppercase(); if !upper.starts_with("SET ") { return false; @@ -690,6 +778,10 @@ fn is_mysql_session_restore_statement(statement: &str) -> bool { } fn leading_executable_sql(sql: &str) -> &str { + leading_executable_sql_with_options(sql, SqlParsingOptions::default()) +} + +fn leading_executable_sql_with_options(sql: &str, options: SqlParsingOptions) -> &str { let bytes = sql.as_bytes(); let mut i = 0; @@ -706,6 +798,14 @@ fn leading_executable_sql(sql: &str) -> &str { continue; } + if options.supports_hash_line_comments && bytes[i] == b'#' { + i += 1; + while i < bytes.len() && bytes[i] != b'\n' { + i += 1; + } + continue; + } + if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' { if i + 2 < bytes.len() && (bytes[i + 2] == b'!' || (i + 3 < bytes.len() && &bytes[i + 2..i + 4] == b"M!")) { break; @@ -724,7 +824,7 @@ fn leading_executable_sql(sql: &str) -> &str { &sql[i..] } -fn first_executable_sql_token(sql: &str) -> Option<&str> { +fn first_executable_sql_token_with_options(sql: &str, options: SqlParsingOptions) -> Option<&str> { let bytes = sql.as_bytes(); let mut i = 0; @@ -741,6 +841,14 @@ fn first_executable_sql_token(sql: &str) -> Option<&str> { continue; } + if options.supports_hash_line_comments && bytes[i] == b'#' { + i += 1; + while i < bytes.len() && bytes[i] != b'\n' { + i += 1; + } + continue; + } + if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' { if i + 2 < bytes.len() && (bytes[i + 2] == b'!' || (i + 3 < bytes.len() && &bytes[i + 2..i + 4] == b"M!")) { i += if bytes[i + 2] == b'!' { 3 } else { 4 }; @@ -799,6 +907,25 @@ fn dollar_quote_tag_at(chars: &[char], start: usize) -> Option { } fn has_executable_sql(statement: &str) -> bool { + has_executable_sql_with_options(statement, SqlParsingOptions::default()) +} + +fn executable_sql_slice(statement: &str, options: SqlParsingOptions) -> Option<&str> { + executable_sql_bounds(statement, options).map(|(start, end)| &statement[start..end]) +} + +fn executable_sql_bounds(statement: &str, options: SqlParsingOptions) -> Option<(usize, usize)> { + let trimmed_end = statement.trim_end().len(); + let trimmed = &statement[..trimmed_end]; + let executable = leading_executable_sql_with_options(trimmed, options); + if executable.is_empty() { + return None; + } + let start = trimmed.len() - executable.len(); + Some((start, trimmed_end)) +} + +fn has_executable_sql_with_options(statement: &str, options: SqlParsingOptions) -> bool { let chars = statement.chars().collect::>(); let mut in_line_comment = false; let mut in_block_comment = false; @@ -834,6 +961,13 @@ fn has_executable_sql(statement: &str) -> bool { continue; } + if options.supports_hash_line_comments && ch == '#' { + in_line_comment = true; + previous = Some(ch); + i += 1; + continue; + } + if ch == '/' && next == Some('*') { if is_mysql_executable_comment_start(&chars, i) { return true; @@ -872,8 +1006,9 @@ mod tests { use crate::models::connection::DatabaseType; use super::{ - prepare_sql_file_statement, split_sql_script, starts_with_executable_sql_keyword, SqlFileStatementAction, - SqlStatementSplitter, + find_statement_at_cursor_for_database, prepare_sql_file_statement, split_sql_script, + split_sql_statements_for_database, starts_with_executable_sql_keyword, + starts_with_executable_sql_keyword_for_database, SqlFileStatementAction, SqlStatementSplitter, }; #[test] @@ -979,6 +1114,16 @@ mod tests { assert!(starts_with_executable_sql_keyword("/*M! SELECT 1 */", &["SELECT"])); } + #[test] + fn mysql_hash_comments_are_ignored_for_keyword_detection() { + assert!(starts_with_executable_sql_keyword_for_database( + "# comment only for mysql\nSELECT 1", + &["SELECT"], + DatabaseType::Mysql + )); + assert!(!starts_with_executable_sql_keyword("# comment only for mysql\nSELECT 1", &["SELECT"])); + } + #[test] fn prepares_mysql_executable_comments_for_mysql_compatible_imports() { assert_eq!( @@ -1263,4 +1408,30 @@ SELECT 2;"; assert_eq!(super::find_statement_at_cursor(sql, cursor), "CREATE PROCEDURE foo()\nBEGIN\n SELECT 1;\nEND"); assert_eq!(super::find_statement_at_cursor(sql, next_cursor), "SELECT 2"); } + + #[test] + fn mysql_hash_comments_split_statements_per_issue_428() { + let sql = "SELECT 1; # mysql comment\n\nSELECT 2 # trailing comment"; + assert_eq!( + split_sql_statements_for_database(sql, DatabaseType::Mysql), + vec!["SELECT 1", "SELECT 2 # trailing comment"] + ); + } + + #[test] + fn mysql_current_statement_keeps_inline_hash_comment_per_issue_428() { + let sql = "SELECT 1; # mysql comment\n\nSELECT 2 # trailing comment"; + let cursor = sql[..sql.find("SELECT 2").unwrap()].encode_utf16().count(); + assert_eq!( + find_statement_at_cursor_for_database(sql, cursor, DatabaseType::Mysql), + "SELECT 2 # trailing comment" + ); + } + + #[test] + fn mysql_single_statement_with_inline_comment_stays_executable_per_issue_428() { + let sql = "SELECT 1 # mysql comment"; + let cursor = sql.encode_utf16().count(); + assert_eq!(find_statement_at_cursor_for_database(sql, cursor, DatabaseType::Mysql), "SELECT 1 # mysql comment"); + } } diff --git a/crates/dbx-core/src/sql_dialect.rs b/crates/dbx-core/src/sql_dialect.rs index 34fa38d27..4a12b9c97 100644 --- a/crates/dbx-core/src/sql_dialect.rs +++ b/crates/dbx-core/src/sql_dialect.rs @@ -344,6 +344,7 @@ pub fn is_schema_aware(database_type: DatabaseType) -> bool { | DatabaseType::Trino | DatabaseType::Db2 | DatabaseType::Tdengine + | DatabaseType::DuckDb ) } diff --git a/crates/dbx-web/src/routes/query.rs b/crates/dbx-web/src/routes/query.rs index 5bed8586f..457d1bc00 100644 --- a/crates/dbx-web/src/routes/query.rs +++ b/crates/dbx-web/src/routes/query.rs @@ -62,6 +62,7 @@ pub struct AnalyzeEditableQueryRequest { pub struct FindStatementAtCursorRequest { pub sql: String, pub cursor_pos: usize, + pub database_type: Option, } #[derive(Deserialize)] @@ -338,7 +339,13 @@ pub async fn execute_script( State(state): State>, Json(req): Json, ) -> Result, AppError> { - let statements = dbx_core::sql::split_sql_statements(&req.sql); + let db_type = { + let configs = state.app.configs.read().await; + configs.get(&req.connection_id).map(|config| config.db_type) + }; + let statements = db_type + .map(|db_type| dbx_core::sql::split_sql_statements_for_database(&req.sql, db_type)) + .unwrap_or_else(|| dbx_core::sql::split_sql_statements(&req.sql)); let result = dbx_core::query::execute_statements( &state.app, &req.connection_id, @@ -376,7 +383,11 @@ pub async fn analyze_sql_references( } pub async fn find_statement_at_cursor(Json(req): Json) -> Json { - Json(dbx_core::sql::find_statement_at_cursor(&req.sql, req.cursor_pos)) + Json( + req.database_type + .map(|db_type| dbx_core::sql::find_statement_at_cursor_for_database(&req.sql, req.cursor_pos, db_type)) + .unwrap_or_else(|| dbx_core::sql::find_statement_at_cursor(&req.sql, req.cursor_pos)), + ) } pub async fn prepare_query_pagination_execution_plan( diff --git a/src-tauri/src/commands/query.rs b/src-tauri/src/commands/query.rs index 1dc371d1a..e9cf97f25 100644 --- a/src-tauri/src/commands/query.rs +++ b/src-tauri/src/commands/query.rs @@ -3,6 +3,7 @@ use tauri::State; use crate::commands::connection::AppState; use dbx_core::db; +use dbx_core::models::connection::DatabaseType; use dbx_core::sql::split_sql_statements; // Re-export core functions for use by other modules (e.g., sql_file.rs) @@ -119,11 +120,19 @@ pub async fn execute_script( sql: String, schema: Option, ) -> Result { + let db_type = { + let configs = state.configs.read().await; + configs.get(&connection_id).map(|config| config.db_type) + }; + dbx_core::query::execute_statements( &state, &connection_id, &database, - &split_sql_statements(&sql), + &db_type.map_or_else( + || split_sql_statements(&sql), + |db_type| dbx_core::sql::split_sql_statements_for_database(&sql, db_type), + ), schema.as_deref(), ) .await @@ -156,8 +165,14 @@ pub async fn analyze_sql_references( } #[tauri::command] -pub fn find_statement_at_cursor(sql: String, cursor_pos: usize) -> Result { - Ok(dbx_core::sql::find_statement_at_cursor(&sql, cursor_pos)) +pub fn find_statement_at_cursor( + sql: String, + cursor_pos: usize, + database_type: Option, +) -> Result { + Ok(database_type + .map(|db_type| dbx_core::sql::find_statement_at_cursor_for_database(&sql, cursor_pos, db_type)) + .unwrap_or_else(|| dbx_core::sql::find_statement_at_cursor(&sql, cursor_pos))) } #[tauri::command]