fix(sqlserver): report affected rows across SQL batches
This commit is contained in:
parent
9440835c94
commit
2c91cdaba3
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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::<u64>(), 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore = "requires DBX_TEST_SQLSERVER_HOST and DBX_TEST_SQLSERVER_PASSWORD"]
|
||||
async fn sqlserver_single_result_query_drains_later_results() {
|
||||
|
|
|
|||
Loading…
Reference in New Issue