diff --git a/crates/dbx-core/src/db/mysql.rs b/crates/dbx-core/src/db/mysql.rs index 01bf8be64..ffc6d1931 100644 --- a/crates/dbx-core/src/db/mysql.rs +++ b/crates/dbx-core/src/db/mysql.rs @@ -3676,6 +3676,91 @@ async fn execute_result_set_with_text_protocol_on_conn( }) } +async fn execute_result_sets_with_text_protocol_on_conn( + conn: &mut mysql_async::Conn, + sql: &str, + row_limit: usize, + max_rows: Option, + start: Instant, +) -> Result, String> { + let mut result = conn.query_iter(sql).await.map_err(|e| e.to_string())?; + let mut results = Vec::new(); + + while advance_to_result_set_with_columns(&mut result).await? { + let columns: Vec = result.columns_ref().iter().map(|c| c.name_str().to_string()).collect(); + let column_types: Vec = result.columns_ref().iter().map(mysql_column_type_name).collect(); + let mut spatial_columns = mysql_spatial_column_builder(result.columns_ref()); + let mut spatial_values = Vec::new(); + let mut truncated = false; + + let rows = if should_collect_text_result_set(sql, row_limit, max_rows) { + let rows: Vec = result.collect().await.map_err(|e| e.to_string())?; + truncated = rows.len() > row_limit; + rows.iter() + .take(row_limit) + .map(|row| { + let (values, srids) = mysql_row_to_json_with_srids(row, &mut spatial_columns); + spatial_values.push(srids); + values + }) + .collect() + } else { + let mut rows = Vec::new(); + let mut stream = result + .stream::() + .await + .map_err(|e| e.to_string())? + .ok_or_else(|| "Empty result set stream".to_string())?; + + while let Some(row) = stream.next().await { + let row = row.map_err(|e| e.to_string())?; + if rows.len() < row_limit { + let (values, srids) = mysql_row_to_json_with_srids(&row, &mut spatial_columns); + rows.push(values); + spatial_values.push(srids); + } else { + truncated = true; + } + } + rows + }; + + results.push(QueryResult { + columns, + column_types, + column_sortables: vec![], + spatial_columns: spatial_columns.finish(), + spatial_values, + rows, + affected_rows: 0, + execution_time_ms: start.elapsed().as_millis(), + truncated, + session_id: None, + has_more: false, + elasticsearch_raw_body: None, + }); + } + + if results.is_empty() { + results.push(QueryResult { + columns: vec![], + column_types: Vec::new(), + column_sortables: vec![], + spatial_columns: vec![], + spatial_values: vec![], + rows: vec![], + affected_rows: result.affected_rows(), + execution_time_ms: start.elapsed().as_millis(), + truncated: false, + session_id: None, + has_more: false, + elasticsearch_raw_body: None, + }); + } + + Ok(results) +} + async fn advance_to_result_set_with_columns( result: &mut mysql_async::QueryResult<'_, '_, mysql_async::TextProtocol>, ) -> Result { @@ -3997,6 +4082,22 @@ pub async fn execute_query_on_conn_with_max_rows( } } +pub async fn execute_query_results_on_conn_with_max_rows( + conn: &mut mysql_async::Conn, + sql: &str, + bare: bool, + max_rows: Option, + dialect: MySqlQueryDialect, +) -> Result, String> { + if is_result_set_query(sql, dialect) && (bare || prefers_text_protocol_query(sql, dialect)) { + let start = Instant::now(); + execute_result_sets_with_text_protocol_on_conn(conn, sql, query_result_row_limit(max_rows), max_rows, start) + .await + } else { + execute_query_on_conn_with_max_rows(conn, sql, bare, max_rows, dialect).await.map(|result| vec![result]) + } +} + fn prefers_text_protocol_query(sql: &str, dialect: MySqlQueryDialect) -> bool { // User-entered result-set queries are not parameterized in DBX. Text protocol // avoids binary result decoding bugs in MySQL-compatible servers and proxies. diff --git a/crates/dbx-core/src/query.rs b/crates/dbx-core/src/query.rs index 79ecc5ee4..0e4248e54 100644 --- a/crates/dbx-core/src/query.rs +++ b/crates/dbx-core/src/query.rs @@ -1909,7 +1909,7 @@ pub async fn execute_multi_core_with_options_for_client_and_progress( } trait MysqlBatchStatementExecutor { - async fn execute_statement(&mut self, statement: &str) -> Result; + async fn execute_statement(&mut self, statement: &str) -> Result, String>; } struct MysqlBatchConnection<'a> { @@ -1922,11 +1922,11 @@ struct MysqlBatchConnection<'a> { } impl MysqlBatchStatementExecutor for MysqlBatchConnection<'_> { - async fn execute_statement(&mut self, statement: &str) -> Result { - wait_for_query_opt( + async fn execute_statement(&mut self, statement: &str) -> Result, String> { + wait_for_result_opt( self.cancel_token.clone(), self.query_timeout, - db::mysql::execute_query_on_conn_with_max_rows( + db::mysql::execute_query_results_on_conn_with_max_rows( &mut *self.conn, statement, self.bare, @@ -1957,9 +1957,15 @@ where } match executor.execute_statement(statement).await { - Ok(result) => { - report_execute_multi_progress(progress, statement_index, statements.len(), &result, true, None); - results.push(ExecuteMultiResult::success_with_index(result, statement_index)); + Ok(statement_results) => { + if let Some(result) = statement_results.last() { + report_execute_multi_progress(progress, statement_index, statements.len(), result, true, None); + } + results.extend( + statement_results + .into_iter() + .map(|result| ExecuteMultiResult::success_with_index(result, statement_index)), + ); } Err(err) => { let action = pool_error_action(db_type, &err); @@ -4100,17 +4106,21 @@ for line in sys.stdin: } struct FakeMysqlBatchExecutor { - outcomes: std::collections::VecDeque>, + outcomes: std::collections::VecDeque, String>>, executed: Vec, } impl MysqlBatchStatementExecutor for FakeMysqlBatchExecutor { - async fn execute_statement(&mut self, statement: &str) -> Result { + async fn execute_statement(&mut self, statement: &str) -> Result, String> { self.executed.push(statement.to_string()); self.outcomes.pop_front().expect("test outcome for statement") } } + fn mysql_batch_result(result: db::QueryResult) -> Result, String> { + Ok(vec![result]) + } + async fn assert_sqlite_batch_error_behavior(failure_first: bool, continue_on_error: bool) { let dir = std::env::temp_dir().join(format!("dbx-query-batch-error-{}", uuid::Uuid::new_v4())); std::fs::create_dir_all(&dir).unwrap(); @@ -4211,9 +4221,9 @@ for line in sys.stdin: let statements = vec!["first".to_string(), "fails".to_string(), "must-not-run".to_string()]; let mut executor = FakeMysqlBatchExecutor { outcomes: std::collections::VecDeque::from([ - Ok(empty_query_result(0)), + mysql_batch_result(empty_query_result(0)), Err("Duplicate entry".to_string()), - Ok(empty_query_result(0)), + mysql_batch_result(empty_query_result(0)), ]), executed: Vec::new(), }; @@ -4234,7 +4244,10 @@ for line in sys.stdin: async fn mysql_batch_reports_progress_for_each_completed_statement() { let statements = vec!["first".to_string(), "fails".to_string(), "must-not-run".to_string()]; let mut executor = FakeMysqlBatchExecutor { - outcomes: std::collections::VecDeque::from([Ok(empty_query_result(0)), Err("Duplicate entry".to_string())]), + outcomes: std::collections::VecDeque::from([ + mysql_batch_result(empty_query_result(0)), + Err("Duplicate entry".to_string()), + ]), executed: Vec::new(), }; let progress_events = Arc::new(std::sync::Mutex::new(Vec::new())); @@ -4281,11 +4294,56 @@ for line in sys.stdin: assert_eq!(error_action, Some(PoolErrorAction::Keep)); } + #[tokio::test] + async fn mysql_batch_preserves_multiple_result_sets_from_one_statement() { + let statements = vec!["CALL testA()".to_string(), "UPDATE users SET active = 1".to_string()]; + let result_set = |value| db::QueryResult { + columns: vec!["value".to_string()], + column_types: vec!["INT".to_string()], + column_sortables: vec![], + spatial_columns: vec![], + spatial_values: vec![], + rows: vec![vec![serde_json::json!(value)]], + affected_rows: 0, + execution_time_ms: 1, + truncated: false, + session_id: None, + has_more: false, + elasticsearch_raw_body: None, + }; + let mut executor = FakeMysqlBatchExecutor { + outcomes: std::collections::VecDeque::from([ + Ok(vec![result_set(1), result_set(2), result_set(3)]), + mysql_batch_result(empty_query_result(1)), + ]), + executed: Vec::new(), + }; + + let (results, error_action) = + execute_mysql_batch_statements(&mut executor, &statements, Some(DatabaseType::Mysql), None, false, None) + .await; + + assert_eq!(executor.executed, statements); + assert_eq!(results.len(), 4); + assert_eq!( + results.iter().map(|result| result.statement_index).collect::>(), + vec![Some(0), Some(0), Some(0), Some(1)] + ); + assert_eq!( + results[..3].iter().map(|result| result.result.rows[0][0].clone()).collect::>(), + vec![serde_json::json!(1), serde_json::json!(2), serde_json::json!(3)] + ); + assert_eq!(error_action, None); + } + #[tokio::test] async fn mysql_batch_stops_when_the_first_statement_fails() { let statements = vec!["fails".to_string(), "must-not-run".to_string()]; let mut executor = FakeMysqlBatchExecutor { - outcomes: std::collections::VecDeque::from([Err("Duplicate entry".to_string()), Ok(empty_query_result(0))]), + outcomes: std::collections::VecDeque::from([ + Err("Duplicate entry".to_string()), + mysql_batch_result(empty_query_result(0)), + ]), executed: Vec::new(), }; @@ -4304,9 +4362,9 @@ for line in sys.stdin: let statements = vec!["first".to_string(), "fails".to_string(), "third".to_string()]; let mut executor = FakeMysqlBatchExecutor { outcomes: std::collections::VecDeque::from([ - Ok(empty_query_result(0)), + mysql_batch_result(empty_query_result(0)), Err("Duplicate entry".to_string()), - Ok(empty_query_result(0)), + mysql_batch_result(empty_query_result(0)), ]), executed: Vec::new(), }; @@ -4329,7 +4387,10 @@ for line in sys.stdin: async fn mysql_batch_continues_when_the_first_statement_fails_and_enabled() { let statements = vec!["fails".to_string(), "second".to_string()]; let mut executor = FakeMysqlBatchExecutor { - outcomes: std::collections::VecDeque::from([Err("Duplicate entry".to_string()), Ok(empty_query_result(0))]), + outcomes: std::collections::VecDeque::from([ + Err("Duplicate entry".to_string()), + mysql_batch_result(empty_query_result(0)), + ]), executed: Vec::new(), }; @@ -4348,9 +4409,9 @@ for line in sys.stdin: let statements = vec!["first".to_string(), "disconnects".to_string(), "must-not-run".to_string()]; let mut executor = FakeMysqlBatchExecutor { outcomes: std::collections::VecDeque::from([ - Ok(empty_query_result(0)), + mysql_batch_result(empty_query_result(0)), Err("connection reset by peer".to_string()), - Ok(empty_query_result(0)), + mysql_batch_result(empty_query_result(0)), ]), executed: Vec::new(), }; diff --git a/crates/dbx-core/tests/live_mysql57.rs b/crates/dbx-core/tests/live_mysql57.rs index c0ee31885..cf89cbcb9 100644 --- a/crates/dbx-core/tests/live_mysql57.rs +++ b/crates/dbx-core/tests/live_mysql57.rs @@ -10,6 +10,7 @@ use dbx_core::sql::{split_sql_statements_for_database, SqlFileRequest}; use dbx_core::sql_file_import::execute_sql_file_path; use dbx_core::storage::Storage; use dbx_core::table_import::parse_xlsx_file; +use mysql_async::prelude::Queryable; use tokio_util::sync::CancellationToken; fn live_mysql_sql_file_config(id: &str) -> ConnectionConfig { @@ -99,6 +100,38 @@ async fn live_mysql57_text_protocol_select_succeeds() { assert_eq!(result.rows, vec![vec![serde_json::json!("1"), serde_json::json!("mysql57")]]); } +#[tokio::test] +#[ignore = "requires a MySQL endpoint that permits stored procedure creation"] +async fn live_mysql_stored_procedure_preserves_all_result_sets() { + let url = std::env::var("DBX_LIVE_MYSQL57_URL").expect("DBX_LIVE_MYSQL57_URL"); + let pool = dbx_core::db::mysql::connect(&url, std::time::Duration::from_secs(5)).await.unwrap(); + let procedure = format!("dbx_issue_4609_{}", uuid::Uuid::new_v4().simple()); + let mut conn = dbx_core::db::mysql::get_conn_with_health_check(&pool).await.unwrap(); + conn.query_drop(format!( + "CREATE PROCEDURE `{procedure}`() BEGIN SELECT 1 AS value; SELECT 2 AS value; SELECT 3 AS value; END" + )) + .await + .unwrap(); + + let results = dbx_core::db::mysql::execute_query_results_on_conn_with_max_rows( + &mut conn, + &format!("CALL `{procedure}`()"), + false, + Some(10), + Default::default(), + ) + .await; + let cleanup = conn.query_drop(format!("DROP PROCEDURE `{procedure}`")).await; + + let results = results.unwrap(); + cleanup.unwrap(); + assert_eq!(results.len(), 3); + assert_eq!( + results.iter().map(|result| result.rows[0][0].clone()).collect::>(), + vec![serde_json::json!("1"), serde_json::json!("2"), serde_json::json!("3")] + ); +} + #[tokio::test] #[ignore = "requires a remote MySQL-compatible endpoint with a limited result-set query"] async fn live_mysql_compatible_limited_text_protocol_query_succeeds() {