mirror of
https://github.com/t8y2/dbx.git
synced 2026-10-02 02:34:42 +08:00
fix(postgres): accept safe psql control commands preserving stop-on-error
This commit is contained in:
@@ -3789,7 +3789,13 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_gaussdb_on_error_stop_overrides_continue_on_error_at_script_position() {
|
||||
async fn streaming_postgres_family_on_error_stop_overrides_continue_on_error_at_script_position() {
|
||||
for db_type in [DatabaseType::Postgres, DatabaseType::OpenGauss, DatabaseType::Gaussdb] {
|
||||
assert_streaming_psql_on_error_stop_at_script_position(db_type).await;
|
||||
}
|
||||
}
|
||||
|
||||
async fn assert_streaming_psql_on_error_stop_at_script_position(db_type: DatabaseType) {
|
||||
let dir = std::env::temp_dir().join(format!("dbx-sql-file-stop-on-error-{}", uuid::Uuid::new_v4()));
|
||||
std::fs::create_dir_all(&dir).unwrap();
|
||||
let storage = crate::persistence::test_storage::open(&dir.join("storage.db")).await.unwrap();
|
||||
@@ -3797,7 +3803,7 @@ mod tests {
|
||||
let config: crate::models::connection::ConnectionConfig = serde_json::from_value(serde_json::json!({
|
||||
"id": "gauss-stream",
|
||||
"name": "GaussDB stream test",
|
||||
"db_type": "gaussdb",
|
||||
"db_type": db_type,
|
||||
"host": "localhost",
|
||||
"port": 5432,
|
||||
"username": "",
|
||||
@@ -3815,7 +3821,7 @@ mod tests {
|
||||
.await;
|
||||
|
||||
let path = temporary_sql_file(
|
||||
b"CREATE TABLE side_effects(value INTEGER);\nINSERT INTO missing_before_control VALUES (1);\nINSERT INTO side_effects VALUES (1);\n\\set ON_ERROR_STOP on\nINSERT INTO missing_after_control VALUES (1);\nINSERT INTO side_effects VALUES (2);",
|
||||
b"CREATE TABLE side_effects(value INTEGER);\nINSERT INTO missing_before_control VALUES (1);\nINSERT INTO side_effects VALUES (1);\n\\set ON_ERROR_STOP on\n\\echo progress; -- display text\nINSERT INTO missing_after_control VALUES (1);\nINSERT INTO side_effects VALUES (2);",
|
||||
)
|
||||
.await;
|
||||
let request = SqlFileRequest {
|
||||
|
||||
@@ -8995,7 +8995,13 @@ for line in sys.stdin:
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gaussdb_on_error_stop_overrides_continue_on_error() {
|
||||
async fn postgres_family_on_error_stop_overrides_continue_on_error() {
|
||||
for db_type in [DatabaseType::Postgres, DatabaseType::OpenGauss, DatabaseType::Gaussdb] {
|
||||
assert_psql_on_error_stop_overrides_continue_on_error(db_type).await;
|
||||
}
|
||||
}
|
||||
|
||||
async fn assert_psql_on_error_stop_overrides_continue_on_error(db_type: DatabaseType) {
|
||||
let dir = std::env::temp_dir().join(format!("dbx-query-gaussdb-on-error-stop-{}", uuid::Uuid::new_v4()));
|
||||
std::fs::create_dir_all(&dir).unwrap();
|
||||
let storage = crate::persistence::test_storage::open(&dir.join("storage.db")).await.unwrap();
|
||||
@@ -9007,7 +9013,7 @@ for line in sys.stdin:
|
||||
connections.insert(connection_id.to_string(), PoolKind::Sqlite(sqlite));
|
||||
})
|
||||
.await;
|
||||
state.configs.write().await.insert(connection_id.to_string(), test_connection_config(DatabaseType::Gaussdb));
|
||||
state.configs.write().await.insert(connection_id.to_string(), test_connection_config(db_type));
|
||||
|
||||
let results = execute_multi_core_with_options(
|
||||
&state,
|
||||
|
||||
+153
-23
@@ -232,6 +232,7 @@ impl SqlDialectProfile {
|
||||
Self {
|
||||
supports_oracle_style_routine_bodies: true,
|
||||
supports_slash_line_block_delimiter: true,
|
||||
supports_psql_control_commands: true,
|
||||
..Self::default()
|
||||
}
|
||||
}
|
||||
@@ -415,6 +416,7 @@ pub struct SqlStatementSplitter {
|
||||
pending_mysql_line_comment_dashes: bool,
|
||||
custom_delimiter: Option<String>,
|
||||
stop_on_error: bool,
|
||||
pending_psql_command: Option<String>,
|
||||
options: SqlParsingOptions,
|
||||
}
|
||||
|
||||
@@ -446,6 +448,32 @@ impl SqlStatementSplitter {
|
||||
}
|
||||
|
||||
while i < chars.len() {
|
||||
if let Some(command) = &mut self.pending_psql_command {
|
||||
let ch = chars[i];
|
||||
command.push(ch);
|
||||
if ch == '\n' {
|
||||
self.finish_psql_command(&mut statements);
|
||||
}
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
// A psql command consumes the whole line, including quotes, comment
|
||||
// markers and semicolons in display text. Buffer it across chunks
|
||||
// before letting the SQL scanner interpret those characters.
|
||||
if self.options.profile.supports_psql_control_commands
|
||||
&& chars[i] == '\\'
|
||||
&& !self.in_single_quote
|
||||
&& !self.in_double_quote
|
||||
&& !self.in_backtick
|
||||
&& !self.in_line_comment
|
||||
&& !self.in_block_comment
|
||||
&& self.dollar_quote_tag.is_none()
|
||||
&& self.buffer.rsplit('\n').next().unwrap_or("").trim().is_empty()
|
||||
{
|
||||
self.pending_psql_command = Some(String::from("\\"));
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
if let Some(tag) = &self.dollar_quote_tag {
|
||||
let tag_chars = tag.chars().collect::<Vec<_>>();
|
||||
if starts_with_chars(&chars, i, &tag_chars) {
|
||||
@@ -605,17 +633,6 @@ impl SqlStatementSplitter {
|
||||
let buf_end = self.buffer.len() - 1;
|
||||
let last_line_start = self.buffer[..buf_end].rfind('\n').map_or(0, |p| p + 1);
|
||||
let last_line = self.buffer[last_line_start..buf_end].trim();
|
||||
if self.options.profile.supports_psql_control_commands {
|
||||
if let Some(command) = parse_ignorable_psql_control_command(last_line) {
|
||||
if command == PsqlControlCommand::OnErrorStop {
|
||||
self.stop_on_error = true;
|
||||
}
|
||||
self.buffer.truncate(last_line_start);
|
||||
self.previous = self.buffer.chars().next_back();
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
if self.options.profile.supports_slash_line_block_delimiter && last_line == "/" {
|
||||
let before = self.buffer[..last_line_start].trim();
|
||||
if has_executable_sql_with_options(before, self.options) {
|
||||
@@ -673,29 +690,34 @@ impl SqlStatementSplitter {
|
||||
statements
|
||||
}
|
||||
|
||||
fn finish_psql_command(&mut self, statements: &mut Vec<SqlStatementWithControl>) {
|
||||
let command = self.pending_psql_command.take().expect("pending psql command");
|
||||
if let Some(control) = parse_ignorable_psql_control_command(command.trim()) {
|
||||
self.stop_on_error |= control == PsqlControlCommand::OnErrorStop;
|
||||
} else {
|
||||
// Unsupported commands must reach the SQL executor and fail visibly;
|
||||
// do not silently discard variable assignments or execution commands.
|
||||
self.options.profile.supports_psql_control_commands = false;
|
||||
statements.extend(self.push_chunk_with_control(&command));
|
||||
self.options.profile.supports_psql_control_commands = true;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn finish(self) -> Vec<String> {
|
||||
self.finish_with_control().into_iter().map(|statement| statement.sql).collect()
|
||||
}
|
||||
|
||||
pub fn finish_with_control(mut self) -> Vec<SqlStatementWithControl> {
|
||||
let mut statements = Vec::new();
|
||||
if self.pending_psql_command.is_some() {
|
||||
self.finish_psql_command(&mut statements);
|
||||
}
|
||||
if self.pending_mysql_line_comment_dashes {
|
||||
self.in_line_comment = true;
|
||||
self.pending_mysql_line_comment_dashes = false;
|
||||
}
|
||||
let trimmed = self.buffer.trim();
|
||||
let last_line = trimmed.rsplit('\n').next().unwrap_or(trimmed).trim();
|
||||
if self.options.profile.supports_psql_control_commands {
|
||||
if let Some(command) = parse_ignorable_psql_control_command(last_line) {
|
||||
if command == PsqlControlCommand::OnErrorStop {
|
||||
self.stop_on_error = true;
|
||||
}
|
||||
let line_start = self.buffer.rfind('\n').map_or(0, |pos| pos + 1);
|
||||
self.buffer.truncate(line_start);
|
||||
}
|
||||
}
|
||||
let trimmed = self.buffer.trim();
|
||||
let last_line = trimmed.rsplit('\n').next().unwrap_or(trimmed).trim();
|
||||
let before_last_line = trimmed.rsplit_once('\n').map(|x| x.0).unwrap_or("").trim();
|
||||
if self.options.profile.supports_custom_delimiter_commands
|
||||
&& parse_delimiter_command(last_line)
|
||||
@@ -1686,6 +1708,13 @@ fn parse_ignorable_psql_control_command(line: &str) -> Option<PsqlControlCommand
|
||||
let mut parts = line.split_whitespace();
|
||||
let command = parts.next()?;
|
||||
|
||||
if command.eq_ignore_ascii_case("\\echo") {
|
||||
// Display-only output can be omitted, but another meta-command on the
|
||||
// same line or psql's backtick command substitution must not be ignored.
|
||||
let arguments = &line[line.find(command)? + command.len()..];
|
||||
return (!arguments.contains(['\\', '`'])).then_some(PsqlControlCommand::ClientDisplay);
|
||||
}
|
||||
|
||||
if command.eq_ignore_ascii_case("\\timing") {
|
||||
let valid =
|
||||
parts.next().is_none_or(|value| value.eq_ignore_ascii_case("on") || value.eq_ignore_ascii_case("off"));
|
||||
@@ -1723,7 +1752,15 @@ fn preprocess_psql_control_commands(sql: &str, profile: SqlDialectProfile) -> (S
|
||||
let command = (!scanner.is_masked()).then(|| parse_ignorable_psql_control_command(line.trim())).flatten();
|
||||
if let Some(command) = command {
|
||||
stop_on_error |= command == PsqlControlCommand::OnErrorStop;
|
||||
output.extend(line.chars().map(|ch| if ch.is_whitespace() { ch } else { ' ' }));
|
||||
// Cursor positions arrive as UTF-16 offsets from the editor. Keep
|
||||
// those offsets stable even when display text contains emoji.
|
||||
for ch in line.chars() {
|
||||
if ch.is_whitespace() {
|
||||
output.push(ch);
|
||||
} else {
|
||||
output.extend(std::iter::repeat_n(' ', ch.len_utf16()));
|
||||
}
|
||||
}
|
||||
if segment.ends_with('\n') {
|
||||
output.push('\n');
|
||||
scanner.step(sql, offset + line.len(), '\n');
|
||||
@@ -4360,6 +4397,99 @@ END;";
|
||||
assert_eq!(plan.statements, vec!["SELECT 1", "SELECT 2"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn postgres_family_accepts_display_controls_and_preserves_on_error_stop() {
|
||||
let sql = "\\set ON_ERROR_STOP on\r\n\\echo '开始; SELECT 99; -- output'\r\n\\timing on\n\\set VERBOSITY verbose\nSELECT 1;\n\\echo done; SELECT 99; /* output */\nSELECT 2;\n\\echo";
|
||||
for db_type in [DatabaseType::Postgres, DatabaseType::OpenGauss, DatabaseType::Gaussdb] {
|
||||
let plan = super::sql_execution_plan_for_database(sql, db_type);
|
||||
assert_eq!(plan.statements, vec!["SELECT 1", "SELECT 2"], "{db_type:?}");
|
||||
assert!(plan.stop_on_error);
|
||||
// Import parsing must agree regardless of where input chunks end.
|
||||
for chunk_size in 1..=sql.len() {
|
||||
let mut splitter = SqlStatementSplitter::with_options(SqlParsingOptions::for_database_type(db_type));
|
||||
let chars = sql.chars().collect::<Vec<_>>();
|
||||
let mut statements = Vec::new();
|
||||
for chunk in chars.chunks(chunk_size) {
|
||||
statements.extend(splitter.push_chunk_with_control(&chunk.iter().collect::<String>()));
|
||||
}
|
||||
statements.extend(splitter.finish_with_control());
|
||||
assert_eq!(statements.iter().map(|s| s.sql.as_str()).collect::<Vec<_>>(), vec!["SELECT 1", "SELECT 2"]);
|
||||
assert!(statements.iter().all(|s| s.stop_on_error));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn postgres_psql_controls_do_not_execute_unfinished_sql_buffer() {
|
||||
let sql = "SELECT\n\\echo progress; -- text\n 1;\nSELECT 2;";
|
||||
let plan = super::sql_execution_plan_for_database(sql, DatabaseType::Postgres);
|
||||
assert_eq!(plan.statements.len(), 2);
|
||||
assert_eq!(plan.statements[0].split_whitespace().collect::<Vec<_>>(), vec!["SELECT", "1"]);
|
||||
assert!(!plan.stop_on_error);
|
||||
let mut splitter =
|
||||
SqlStatementSplitter::with_options(SqlParsingOptions::for_database_type(DatabaseType::Postgres));
|
||||
assert!(splitter.push_chunk("SELECT\n\\echo progress; --").is_empty());
|
||||
assert_eq!(splitter.push_chunk(" text\n 1;\nSELECT 2;"), vec!["SELECT\n 1", "SELECT 2"]);
|
||||
assert!(splitter.finish().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn postgres_psql_controls_inside_sql_literals_and_comments_are_preserved() {
|
||||
for sql in [
|
||||
"SELECT 'before\n\\echo literal\nafter';",
|
||||
"SELECT \"before\n\\echo literal\nafter\";",
|
||||
"SELECT $$before\n\\set ON_ERROR_STOP on\nafter$$;",
|
||||
"SELECT $body$before\n\\echo literal\nafter$body$;",
|
||||
"/* before\n\\echo comment\nafter */\nSELECT 1;",
|
||||
"-- \\echo comment\nSELECT 1;",
|
||||
"SELECT 'unfinished\n\\echo literal",
|
||||
"SELECT $$unfinished\n\\set ON_ERROR_STOP on",
|
||||
] {
|
||||
let plan = super::sql_execution_plan_for_database(sql, DatabaseType::Postgres);
|
||||
assert!(!plan.stop_on_error);
|
||||
assert_eq!(plan.statements, vec![sql.trim_end_matches(';')]);
|
||||
let mut splitter =
|
||||
SqlStatementSplitter::with_options(SqlParsingOptions::for_database_type(DatabaseType::Postgres));
|
||||
let mut statements = splitter.push_chunk(sql);
|
||||
statements.extend(splitter.finish());
|
||||
assert_eq!(statements, plan.statements);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn postgres_psql_display_text_preserves_editor_cursor_offsets() {
|
||||
let sql = "\\echo 🚀 开始\nSELECT 1;\nSELECT 2;";
|
||||
let cursor = sql[..sql.find("SELECT 2").unwrap() + "SELECT ".len()].encode_utf16().count();
|
||||
assert_eq!(find_statement_at_cursor_for_database(sql, cursor, DatabaseType::Postgres), "SELECT 2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn postgres_does_not_silently_ignore_variable_or_execution_commands() {
|
||||
for command in [
|
||||
"\\set target_schema private",
|
||||
"\\set ON_ERROR_STOP off",
|
||||
"\\set ON_ERROR_STOP on extra",
|
||||
"\\echo done \\\\ SELECT 99",
|
||||
"\\echo `touch should_not_run`",
|
||||
"\\copy users FROM 'users.csv'",
|
||||
"\\gexec",
|
||||
"\\i script.sql",
|
||||
"\\! echo shell",
|
||||
] {
|
||||
let sql = format!("{command}\nSELECT 1;");
|
||||
let plan = super::sql_execution_plan_for_database(&sql, DatabaseType::Postgres);
|
||||
assert!(plan.statements[0].starts_with(command), "{command}");
|
||||
assert!(!plan.stop_on_error);
|
||||
let mut splitter =
|
||||
SqlStatementSplitter::with_options(SqlParsingOptions::for_database_type(DatabaseType::Postgres));
|
||||
let mut statements = splitter.push_chunk(&sql);
|
||||
statements.extend(splitter.finish());
|
||||
assert!(statements[0].starts_with(command), "{command}");
|
||||
}
|
||||
let sql = "\\echo message\nSELECT 1;";
|
||||
assert!(split_sql_statements_for_database(sql, DatabaseType::Mysql)[0].starts_with("\\echo"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gaussdb_split_separates_dollar_quoted_function_from_following_statements() {
|
||||
let sql = "\
|
||||
|
||||
Reference in New Issue
Block a user