fix(postgres): accept safe psql control commands preserving stop-on-error

This commit is contained in:
zipg
2026-10-01 23:41:23 +08:00
committed by GitHub
parent 192313d33f
commit 788ac11f9b
3 changed files with 170 additions and 28 deletions
+9 -3
View File
@@ -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 {
+8 -2
View File
@@ -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
View File
@@ -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 = "\