From 2c91cdaba32b7ca68b446363c40175454e659db9 Mon Sep 17 00:00:00 2001 From: zipg Date: Wed, 5 Aug 2026 09:45:38 +0800 Subject: [PATCH] fix(sqlserver): report affected rows across SQL batches --- crates/dbx-core/src/db/sqlserver.rs | 85 +++++++++++++++++-- .../dbx-core/tests/sqlserver_batch_results.rs | 80 ++++++++++++++++- 2 files changed, 159 insertions(+), 6 deletions(-) diff --git a/crates/dbx-core/src/db/sqlserver.rs b/crates/dbx-core/src/db/sqlserver.rs index 1b2c939fc..e8e599c80 100644 --- a/crates/dbx-core/src/db/sqlserver.rs +++ b/crates/dbx-core/src/db/sqlserver.rs @@ -2699,15 +2699,53 @@ fn is_dbx_sqlserver_row_number_page_sql(sql: &str) -> bool { } fn sqlserver_batch_can_use_execute(sql: &str) -> bool { + let contains_transaction_control = contains_transaction_control(sql); !requires_simple_query_batch(sql) // Tiberius executes this path through an RPC. SQL Server rejects an RPC // that changes @@TRANCOUNT with error 266, while a regular batch is allowed // to leave an explicit transaction open for a later COMMIT or ROLLBACK. - && !contains_transaction_control(sql) + && (!contains_transaction_control || sqlserver_balanced_transaction_batch_can_use_execute(sql)) && !sqlserver_batch_may_return_result_set(sql) && !sqlserver_dml_output_returns_rows(sql) } +fn sqlserver_balanced_transaction_batch_can_use_execute(sql: &str) -> bool { + let tokens = top_level_sqlserver_tokens(sql); + if tokens.iter().any(|token| { + matches!(token.text.as_str(), "IF" | "ELSE" | "WHILE" | "TRY" | "CATCH" | "GOTO" | "RETURN" | "ROLLBACK") + }) { + return false; + } + + let mut transaction_depth = 0usize; + let mut saw_begin = false; + let mut saw_commit = false; + + for (index, token) in tokens.iter().enumerate() { + if token.text == "SET" && tokens.get(index + 1).is_some_and(|next| next.text == "IMPLICIT_TRANSACTIONS") { + return false; + } + + if token.text == "BEGIN" { + if tokens.get(index + 1).is_some_and(|next| next.text == "DISTRIBUTED") { + return false; + } + if tokens.get(index + 1).is_some_and(|next| matches!(next.text.as_str(), "TRANSACTION" | "TRAN")) { + transaction_depth += 1; + saw_begin = true; + } + } else if token.text == "COMMIT" { + if transaction_depth == 0 { + return false; + } + transaction_depth -= 1; + saw_commit = true; + } + } + + saw_begin && saw_commit && transaction_depth == 0 +} + fn sqlserver_batch_may_return_result_set(sql: &str) -> bool { crate::sql::split_sql_statements(sql).iter().any(|statement| { let tokens = top_level_sqlserver_tokens(statement); @@ -2717,10 +2755,11 @@ fn sqlserver_batch_may_return_result_set(sql: &str) -> bool { } let starts_with_cte_dml = tokens.first().is_some_and(|token| token.text == "WITH") && tokens.iter().any(|token| matches!(token.text.as_str(), "INSERT" | "UPDATE" | "DELETE" | "MERGE")); - tokens.iter().any(|token| { - matches!(token.text.as_str(), "SELECT" | "EXEC" | "EXECUTE" | "TABLE") - || (token.text == "WITH" && !starts_with_cte_dml) - }) + tokens.first().is_some_and(|token| token.text == "TABLE") + || tokens.iter().any(|token| { + matches!(token.text.as_str(), "SELECT" | "EXEC" | "EXECUTE") + || (token.text == "WITH" && !starts_with_cte_dml) + }) }) } @@ -3085,15 +3124,51 @@ mod tests { assert!(sqlserver_batch_can_use_execute( "INSERT INTO dbo.user_archive(id, name) SELECT id, name FROM dbo.users WHERE active = 0;" )); + assert!(sqlserver_batch_can_use_execute( + "DECLARE @target TABLE (id INT); INSERT INTO @target(id) SELECT id FROM (VALUES (1), (2), (3)) AS source(id);" + )); assert!(!sqlserver_batch_can_use_execute( "INSERT INTO dbo.user_archive(id) SELECT id FROM dbo.users; SELECT COUNT(*) FROM dbo.user_archive;" )); + assert!(!sqlserver_batch_can_use_execute( + "DECLARE @target TABLE (id INT); INSERT INTO @target(id) VALUES (1); SELECT id FROM @target;" + )); assert!(sqlserver_batch_can_use_execute("DELETE FROM dbo.users WHERE id = 1;")); assert!(sqlserver_batch_can_use_execute( "MERGE dbo.t AS t USING dbo.s AS s ON t.id = s.id WHEN MATCHED THEN UPDATE SET name = s.name;" )); } + #[test] + fn sqlserver_balanced_transaction_dml_batches_use_execute_for_affected_rows() { + assert!(sqlserver_batch_can_use_execute( + "BEGIN TRANSACTION; UPDATE dbo.users SET active = 0 WHERE id = 1; COMMIT TRANSACTION;" + )); + assert!(sqlserver_batch_can_use_execute( + "BEGIN TRAN; UPDATE dbo.users SET active = 0 WHERE id = 1; DELETE dbo.audit WHERE user_id = 1; COMMIT;" + )); + assert!(sqlserver_batch_can_use_execute( + "BEGIN TRAN outer_tx; BEGIN TRAN inner_tx; UPDATE dbo.users SET active = 0; COMMIT TRAN inner_tx; COMMIT TRAN outer_tx;" + )); + + assert!(!sqlserver_batch_can_use_execute("BEGIN TRAN; UPDATE dbo.users SET active = 0;")); + assert!(!sqlserver_batch_can_use_execute("UPDATE dbo.users SET active = 0; COMMIT TRAN;")); + assert!(!sqlserver_batch_can_use_execute("BEGIN TRAN; UPDATE dbo.users SET active = 0; ROLLBACK TRAN;")); + assert!(!sqlserver_batch_can_use_execute("BEGIN TRAN; IF @should_commit = 1 COMMIT TRAN;")); + assert!(!sqlserver_batch_can_use_execute( + "BEGIN TRY; BEGIN TRAN; UPDATE dbo.users SET active = 0; COMMIT TRAN; END TRY; BEGIN CATCH; ROLLBACK TRAN; END CATCH;" + )); + assert!(!sqlserver_batch_can_use_execute( + "BEGIN DISTRIBUTED TRAN; UPDATE dbo.users SET active = 0; COMMIT TRAN;" + )); + assert!(!sqlserver_batch_can_use_execute( + "SET IMPLICIT_TRANSACTIONS ON; UPDATE dbo.users SET active = 0; COMMIT TRAN;" + )); + assert!(!sqlserver_batch_can_use_execute( + "BEGIN TRAN; UPDATE dbo.users SET active = 0; COMMIT TRAN; SELECT 1;" + )); + } + #[test] fn sqlserver_cte_dml_batches_use_execute_for_affected_rows() { assert!(sqlserver_batch_can_use_execute( diff --git a/crates/dbx-core/tests/sqlserver_batch_results.rs b/crates/dbx-core/tests/sqlserver_batch_results.rs index 0ce20094e..a6bb4a9b0 100644 --- a/crates/dbx-core/tests/sqlserver_batch_results.rs +++ b/crates/dbx-core/tests/sqlserver_batch_results.rs @@ -8,7 +8,8 @@ async fn connect_sqlserver() -> sqlserver::SqlServerClient { let port = std::env::var("DBX_TEST_SQLSERVER_PORT").ok().and_then(|value| value.parse().ok()).unwrap_or(1433); let user = std::env::var("DBX_TEST_SQLSERVER_USER").unwrap_or_else(|_| "sa".to_string()); let password = std::env::var("DBX_TEST_SQLSERVER_PASSWORD").expect("DBX_TEST_SQLSERVER_PASSWORD"); - sqlserver::connect_with_port_explicit(&host, port, true, &user, &password, Some("master"), Duration::from_secs(15)) + let database = std::env::var("DBX_TEST_SQLSERVER_DATABASE").unwrap_or_else(|_| "master".to_string()); + sqlserver::connect_with_port_explicit(&host, port, true, &user, &password, Some(&database), Duration::from_secs(15)) .await .expect("connect to SQL Server") } @@ -53,6 +54,83 @@ async fn sqlserver_insert_select_reports_affected_rows() { assert!(result.rows.is_empty()); } +#[tokio::test] +#[ignore = "requires DBX_TEST_SQLSERVER_HOST and DBX_TEST_SQLSERVER_PASSWORD"] +async fn sqlserver_transaction_batch_reports_affected_rows() { + let mut client = connect_sqlserver().await; + sqlserver::execute_simple_batch_with_max_rows( + &mut client, + "CREATE TABLE #dbx_transaction_rows (id INT NOT NULL PRIMARY KEY, value INT NOT NULL); \ + INSERT INTO #dbx_transaction_rows (id, value) VALUES (1, 0), (2, 0), (3, 0), (4, 0);", + None, + ) + .await + .expect("create transaction row-count fixture"); + + let plain = sqlserver::execute_query( + &mut client, + "UPDATE #dbx_transaction_rows SET value = value + 1 WHERE id IN (1, 2, 3);", + ) + .await + .expect("execute plain UPDATE"); + assert_eq!(plain.affected_rows, 3); + + let committed = sqlserver::execute_query( + &mut client, + "BEGIN TRANSACTION; \ + UPDATE #dbx_transaction_rows SET value = value + 1 WHERE id IN (1, 2, 3); \ + COMMIT TRANSACTION;", + ) + .await + .expect("execute committed transaction UPDATE"); + assert_eq!(committed.affected_rows, 3); + + let persisted = + sqlserver::execute_query(&mut client, "SELECT SUM(value) AS total_value FROM #dbx_transaction_rows;") + .await + .expect("read committed transaction values"); + assert_eq!(persisted.rows[0][0].as_i64(), Some(6)); + + let multiple = sqlserver::execute_query( + &mut client, + "BEGIN TRANSACTION; \ + UPDATE #dbx_transaction_rows SET value = value + 1 WHERE id IN (1, 2); \ + DELETE FROM #dbx_transaction_rows WHERE id = 4; \ + COMMIT TRANSACTION;", + ) + .await + .expect("execute multiple DML statements in a committed transaction"); + assert_eq!(multiple.affected_rows, 3); +} + +#[tokio::test] +#[ignore = "requires DBX_TEST_SQLSERVER_HOST and DBX_TEST_SQLSERVER_PASSWORD"] +async fn sqlserver_driver_execute_reports_balanced_transaction_rows() { + let mut client = connect_sqlserver().await; + client + .simple_query( + "CREATE TABLE #dbx_driver_transaction_rows (id INT NOT NULL PRIMARY KEY, value INT NOT NULL); \ + INSERT INTO #dbx_driver_transaction_rows (id, value) VALUES (1, 0), (2, 0), (3, 0), (4, 0);", + ) + .await + .expect("create driver transaction row-count fixture") + .into_results() + .await + .expect("drain fixture setup results"); + + let result = client + .execute( + "BEGIN TRANSACTION; \ + UPDATE #dbx_driver_transaction_rows SET value = value + 1 WHERE id IN (1, 2, 3); \ + COMMIT TRANSACTION;", + &[], + ) + .await + .expect("execute balanced transaction through RPC"); + + assert_eq!(result.rows_affected().iter().sum::(), 3); +} + #[tokio::test] #[ignore = "requires DBX_TEST_SQLSERVER_HOST and DBX_TEST_SQLSERVER_PASSWORD"] async fn sqlserver_single_result_query_drains_later_results() {