From 1b3137d9474f459443a58df0bdbcb20577afe1ca Mon Sep 17 00:00:00 2001 From: Syngnat Date: Tue, 14 Jul 2026 12:22:07 +0800 Subject: [PATCH] =?UTF-8?q?=F0=9F=90=9B=20fix(transaction):=20=E4=BF=AE?= =?UTF-8?q?=E5=A4=8D=E5=8C=BF=E5=90=8DSQL=E5=9D=97=E8=87=AA=E5=8A=A8?= =?UTF-8?q?=E6=8F=90=E4=BA=A4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 区分匿名 BEGIN...END 块与显式事务控制语句 - 统一前后端方言判定并挂起 Oracle、达梦和 SQL Server 块内 DML - 补充编辑器、事务控制器及 SQL Server 回滚链路回归测试 --- .../QueryEditor.external-sql-save.test.tsx | 41 +++++++--- frontend/src/components/QueryEditor.tsx | 4 +- .../ai/aiSqlEditorTransactionInsights.ts | 9 +-- .../src/utils/sqlEditorTransaction.test.ts | 30 +++++++ frontend/src/utils/sqlEditorTransaction.ts | 78 ++++++++++++++++++- internal/app/methods_db_multi_test.go | 54 +++++++++++++ internal/app/methods_db_transaction.go | 33 ++++---- internal/app/methods_db_transaction_test.go | 40 ++++++++++ 8 files changed, 254 insertions(+), 35 deletions(-) diff --git a/frontend/src/components/QueryEditor.external-sql-save.test.tsx b/frontend/src/components/QueryEditor.external-sql-save.test.tsx index f5aad0a2..599bce53 100644 --- a/frontend/src/components/QueryEditor.external-sql-save.test.tsx +++ b/frontend/src/components/QueryEditor.external-sql-save.test.tsx @@ -135,7 +135,9 @@ const backendApp = vi.hoisted(() => ({ DBQueryMultiInTransaction: vi.fn(), DBQueryMultiTransactional: vi.fn(), DBCommitTransaction: vi.fn(), + DBCommitTransactionWithTrigger: vi.fn(), DBRollbackTransaction: vi.fn(), + DBRollbackTransactionWithTrigger: vi.fn(), DBGetTables: vi.fn(), DBGetAllColumns: vi.fn(), DBGetDatabases: vi.fn(), @@ -840,7 +842,9 @@ describe('QueryEditor external SQL save', () => { backendApp.DBQueryMultiInTransaction.mockResolvedValue({ success: true, data: [] }); backendApp.DBQueryMultiTransactional.mockResolvedValue({ success: true, data: [] }); backendApp.DBCommitTransaction.mockResolvedValue({ success: true, message: '事务已提交' }); + backendApp.DBCommitTransactionWithTrigger.mockResolvedValue({ success: true, message: '事务已提交' }); backendApp.DBRollbackTransaction.mockResolvedValue({ success: true, message: '事务已回滚' }); + backendApp.DBRollbackTransactionWithTrigger.mockResolvedValue({ success: true, message: '事务已回滚' }); backendApp.DBGetColumns.mockResolvedValue({ success: true, data: [] }); backendApp.DBGetIndexes.mockResolvedValue({ success: true, data: [] }); backendApp.DBGetAllColumns.mockResolvedValue({ success: true, data: [] }); @@ -7536,7 +7540,7 @@ describe('QueryEditor external SQL save', () => { await Promise.resolve(); }); - expect(backendApp.DBCommitTransaction).toHaveBeenCalledWith('tx-1'); + expect(backendApp.DBCommitTransactionWithTrigger).toHaveBeenCalledWith('tx-1', 'manual'); expect(storeState.addSqlLog).toHaveBeenCalledWith(expect.objectContaining({ sql: "START TRANSACTION;\nUPDATE users SET name = 'new' WHERE id = 1;\nCOMMIT;", status: 'success', @@ -7737,7 +7741,7 @@ describe('QueryEditor external SQL save', () => { await Promise.resolve(); }); - expect(backendApp.DBCommitTransaction).toHaveBeenCalledWith('tx-with-dml'); + expect(backendApp.DBCommitTransactionWithTrigger).toHaveBeenCalledWith('tx-with-dml', 'manual'); }); it('shows the pending statement count for multi-SQL manual transactions', async () => { @@ -8198,7 +8202,7 @@ describe('QueryEditor external SQL save', () => { }); expect(textContent(renderer!.root)).toContain('3s 后自动提交'); - expect(backendApp.DBCommitTransaction).not.toHaveBeenCalled(); + expect(backendApp.DBCommitTransactionWithTrigger).not.toHaveBeenCalled(); await act(async () => { vi.advanceTimersByTime(3000); @@ -8206,7 +8210,7 @@ describe('QueryEditor external SQL save', () => { await Promise.resolve(); }); - expect(backendApp.DBCommitTransaction).toHaveBeenCalledWith('tx-auto'); + expect(backendApp.DBCommitTransactionWithTrigger).toHaveBeenCalledWith('tx-auto', 'auto'); expect(backendApp.DBQueryMulti).not.toHaveBeenCalled(); } finally { vi.useRealTimers(); @@ -8246,7 +8250,7 @@ describe('QueryEditor external SQL save', () => { expect(backendApp.DBQueryMulti).not.toHaveBeenCalled(); expect(textContent(renderer!.root)).toContain('自动提交中'); expect(textContent(renderer!.root)).toContain('提交 (1)'); - expect(backendApp.DBCommitTransaction).not.toHaveBeenCalled(); + expect(backendApp.DBCommitTransactionWithTrigger).not.toHaveBeenCalled(); await act(async () => { vi.runOnlyPendingTimers(); @@ -8254,7 +8258,7 @@ describe('QueryEditor external SQL save', () => { await Promise.resolve(); }); - expect(backendApp.DBCommitTransaction).toHaveBeenCalledWith('tx-auto-now'); + expect(backendApp.DBCommitTransactionWithTrigger).toHaveBeenCalledWith('tx-auto-now', 'auto'); expect(textContent(renderer!.root)).not.toContain('自动提交中'); } finally { vi.useRealTimers(); @@ -8815,11 +8819,13 @@ describe('QueryEditor external SQL save', () => { renderer?.unmount(); }); - it('keeps Oracle anonymous PL/SQL blocks intact when running from the editor', async () => { + it('keeps Oracle anonymous PL/SQL block DML pending for a manual transaction', async () => { storeState.connections[0].config.type = 'oracle'; storeState.connections[0].config.database = 'ORCLPDB1'; - backendApp.DBQueryMulti.mockResolvedValueOnce({ + backendApp.DBQueryMultiTransactional.mockResolvedValueOnce({ success: true, + transactionId: 'tx-oracle-block', + transactionPending: true, data: [{ columns: ['affectedRows'], rows: [{ affectedRows: 1 }] }], }); const plsql = [ @@ -8843,11 +8849,28 @@ describe('QueryEditor external SQL save', () => { await Promise.resolve(); }); - expect(backendApp.DBQueryMulti).toHaveBeenCalledWith(expect.anything(), 'ORCLPDB1', plsql, 'query-1'); + expect(backendApp.DBQueryMultiTransactional).toHaveBeenCalledWith(expect.anything(), 'ORCLPDB1', plsql, 'query-1'); + expect(backendApp.DBQueryMulti).not.toHaveBeenCalled(); + expect(storeState.sqlEditorPendingTransactions['tab-1']).toMatchObject({ + id: 'tx-oracle-block', + dbType: 'oracle', + statements: [plsql], + }); + expect(textContent(renderer!.root)).toContain('提交'); + expect(textContent(renderer!.root)).toContain('回滚'); expect(storeState.addSqlLog).toHaveBeenCalledWith(expect.objectContaining({ sql: plsql, status: 'success', })); + + await act(async () => { + await findButton(renderer!, '回滚').props.onClick(); + }); + await act(async () => { + await Promise.resolve(); + await Promise.resolve(); + }); + expect(backendApp.DBRollbackTransactionWithTrigger).toHaveBeenCalledWith('tx-oracle-block', 'manual'); renderer?.unmount(); }); diff --git a/frontend/src/components/QueryEditor.tsx b/frontend/src/components/QueryEditor.tsx index 84cf0cc9..ac76e13d 100644 --- a/frontend/src/components/QueryEditor.tsx +++ b/frontend/src/components/QueryEditor.tsx @@ -6633,13 +6633,13 @@ const QueryEditor: React.FC<{ tab: TabData; isActive?: boolean }> = ({ tab, isAc setActiveResultKey(''); return; } - const useManagedTransaction = shouldUseSqlEditorManagedTransactionForType(connCaps.type, sourceStatements); + const useManagedTransaction = shouldUseSqlEditorManagedTransactionForType(normalizedDbType, sourceStatements); if (useManagedTransaction && pendingSqlTransactionRef.current) { message.warning(translate('query_editor.transaction.message.pending_managed_transaction')); return; } const managedTransactionStatementCount = sourceStatements - .filter((statement) => shouldUseSqlEditorManagedTransactionForType(connCaps.type, [statement])) + .filter((statement) => shouldUseSqlEditorManagedTransactionForType(normalizedDbType, [statement])) .length || sourceStatements.length; const forceReadOnlyResult = connCaps.forceReadOnlyQueryResult; diff --git a/frontend/src/components/ai/aiSqlEditorTransactionInsights.ts b/frontend/src/components/ai/aiSqlEditorTransactionInsights.ts index e8ae9ead..319c5d8c 100644 --- a/frontend/src/components/ai/aiSqlEditorTransactionInsights.ts +++ b/frontend/src/components/ai/aiSqlEditorTransactionInsights.ts @@ -3,6 +3,7 @@ import type { I18nParams } from '../../i18n'; import type { SavedConnection, TabData } from '../../types'; import { findSqlStatementRanges } from '../../utils/sqlStatementSelection'; import { + isSqlEditorTransactionControlStatement, shouldUseSqlEditorManagedTransaction, shouldUseSqlEditorManagedTransactionForType, } from '../../utils/sqlEditorTransaction'; @@ -42,10 +43,6 @@ const splitStatements = (sql: string, dbType = ''): string[] => .map((range) => String(range.text || '').trim()) .filter(Boolean); -const hasTransactionControlStatement = (statement: string): boolean => - /^\s*(begin|commit|rollback|savepoint|release)\b/i.test(statement) - || /^\s*start\s+transaction\b/i.test(statement); - const buildTabSummary = ( tab: TabData | undefined, connections: SavedConnection[], @@ -99,7 +96,7 @@ const buildActiveSqlTabSnapshot = (params: { { oceanBaseProtocol: connection?.config?.oceanBaseProtocol }, ); const statements = splitStatements(sql, dbType); - const hasExplicitTransactionControl = statements.some(hasTransactionControlStatement); + const hasExplicitTransactionControl = statements.some(isSqlEditorTransactionControlStatement); const usesManagedTransaction = shouldUseSqlEditorManagedTransactionForType(dbType, statements); return { @@ -150,7 +147,7 @@ const isRelevantSqlEditorTransactionLog = (log: SqlLog): boolean => { const sql = String(log.sql || ''); const statements = splitStatements(sql); if (shouldUseSqlEditorManagedTransaction(statements)) return true; - if (statements.some(hasTransactionControlStatement)) return true; + if (statements.some(isSqlEditorTransactionControlStatement)) return true; return /\b(transaction|commit|rollback)\b/i.test(sql) || LOG_TRANSACTION_KEYWORD_PATTERN.test(String(log.message || '')); }; diff --git a/frontend/src/utils/sqlEditorTransaction.test.ts b/frontend/src/utils/sqlEditorTransaction.test.ts index fa5c3520..7ebda8a7 100644 --- a/frontend/src/utils/sqlEditorTransaction.test.ts +++ b/frontend/src/utils/sqlEditorTransaction.test.ts @@ -67,6 +67,36 @@ describe('sqlEditorTransaction', () => { ])).toBe(false); }); + it('keeps DML inside anonymous BEGIN...END blocks in a managed transaction', () => { + const sqlServerBlock = [ + 'BEGIN', + " PRINT 'DELETE is text here';", + ' -- INSERT INTO audit_logs(id) VALUES (1);', + " UPDATE users SET name = 'new' WHERE id = 1;", + 'END;', + ].join('\n'); + const oracleBlock = [ + 'BEGIN', + " UPDATE users SET name = 'new' WHERE id = 1;", + 'END;', + ].join('\n'); + + expect(shouldUseSqlEditorManagedTransactionForType( + 'sqlserver', + findSqlStatementRanges(sqlServerBlock, 'sqlserver').map((range) => range.text), + )).toBe(true); + expect(shouldUseSqlEditorManagedTransactionForType( + 'oracle', + findSqlStatementRanges(oracleBlock, 'oracle').map((range) => range.text), + )).toBe(true); + }); + + it('does not wrap BEGIN TRANSACTION as an anonymous block', () => { + expect(shouldUseSqlEditorManagedTransactionForType('sqlserver', [ + "BEGIN TRANSACTION; UPDATE users SET name = 'new' WHERE id = 1; COMMIT TRANSACTION;", + ])).toBe(false); + }); + it.each([ ['trino', 'UPDATE hive.default.orders SET status = \'done\''], ['tdengine', 'INSERT INTO meters(ts, current) VALUES (NOW, 10.2)'], diff --git a/frontend/src/utils/sqlEditorTransaction.ts b/frontend/src/utils/sqlEditorTransaction.ts index 6bc189ad..a2e32d9d 100644 --- a/frontend/src/utils/sqlEditorTransaction.ts +++ b/frontend/src/utils/sqlEditorTransaction.ts @@ -1,6 +1,18 @@ const SQL_EDITOR_DML_KEYWORDS = new Set(['insert', 'update', 'delete', 'replace', 'merge', 'upsert']); const SQL_EDITOR_READ_KEYWORDS = new Set(['select', 'with', 'show', 'describe', 'desc', 'explain', 'pragma', 'values']); const SQL_EDITOR_TRANSACTION_CONTROL_KEYWORDS = new Set(['begin', 'commit', 'rollback', 'savepoint', 'release']); +const SQL_EDITOR_BEGIN_TRANSACTION_CONTROL_KEYWORDS = new Set([ + 'transaction', + 'tran', + 'work', + 'isolation', + 'read', + 'write', + 'deferred', + 'immediate', + 'exclusive', + 'distributed', +]); const SQL_EDITOR_MANAGED_TRANSACTION_UNSUPPORTED_TYPES = new Set([ 'trino', 'tdengine', @@ -253,10 +265,63 @@ const sqlEditorStatementHasManagedWrite = (statement: string): boolean => { return SQL_EDITOR_DML_KEYWORDS.has(leading.keyword); }; -const isSqlEditorTransactionControlStatement = (statement: string): boolean => { - const keyword = readSqlEditorKeyword(String(statement || ''), 0).keyword; - if (SQL_EDITOR_TRANSACTION_CONTROL_KEYWORDS.has(keyword)) return true; - return keyword === 'start' && /\btransaction\b/i.test(statement); +const sqlEditorStatementContainsKeyword = (statement: string, wantedKeyword: string): boolean => { + const text = String(statement || ''); + for (let pos = 0; pos < text.length;) { + const skipped = skipSqlEditorQuotedOrComment(text, pos); + if (skipped !== null) { + pos = skipped; + continue; + } + if (!isSqlEditorKeywordChar(text[pos])) { + pos++; + continue; + } + let end = pos + 1; + while (isSqlEditorKeywordChar(text[end])) { + end++; + } + if (text.slice(pos, end).toLowerCase() === wantedKeyword) { + return true; + } + pos = end; + } + return false; +}; + +const isSqlEditorBeginTransactionControlStatement = (statement: string, keywordEnd: number): boolean => { + const text = String(statement || ''); + const next = skipSqlEditorTrivia(text, keywordEnd); + if (next >= text.length || text[next] === ';') return true; + return SQL_EDITOR_BEGIN_TRANSACTION_CONTROL_KEYWORDS.has(readSqlEditorKeyword(text, keywordEnd).keyword); +}; + +export const isSqlEditorTransactionControlStatement = (statement: string): boolean => { + const text = String(statement || ''); + const leading = readSqlEditorKeyword(text, 0); + if (leading.keyword === 'begin') { + return isSqlEditorBeginTransactionControlStatement(text, leading.end); + } + if (SQL_EDITOR_TRANSACTION_CONTROL_KEYWORDS.has(leading.keyword)) return true; + return leading.keyword === 'start' && readSqlEditorKeyword(text, leading.end).keyword === 'transaction'; +}; + +const isSqlEditorManagedBlockWrite = (type: string, statement: string): boolean => { + const text = String(statement || ''); + const leading = readSqlEditorKeyword(text, 0); + const normalizedType = String(type || '').trim().toLowerCase(); + const isOracleLike = ['oracle', 'dameng', 'dm', 'dm8'].includes(normalizedType); + const isSqlServer = ['sqlserver', 'mssql', 'sql_server', 'sql-server'].includes(normalizedType); + + if (isOracleLike) { + if (leading.keyword !== 'begin' && leading.keyword !== 'declare') return false; + } else if (isSqlServer) { + if (leading.keyword !== 'begin' || isSqlEditorBeginTransactionControlStatement(text, leading.end)) return false; + } else { + return false; + } + + return [...SQL_EDITOR_DML_KEYWORDS].some((keyword) => sqlEditorStatementContainsKeyword(text, keyword)); }; export const shouldUseSqlEditorManagedTransactionForType = ( @@ -271,6 +336,10 @@ export const shouldUseSqlEditorManagedTransactionForType = ( const trimmed = String(statement || '').trim(); if (!trimmed) continue; if (isSqlEditorTransactionControlStatement(trimmed)) return false; + if (isSqlEditorManagedBlockWrite(type, trimmed)) { + hasManagedWrite = true; + continue; + } if (sqlEditorStatementHasManagedWrite(trimmed)) { hasManagedWrite = true; continue; @@ -297,6 +366,7 @@ export const canReusePendingSqlEditorTransactionForType = ( const trimmed = String(statement || '').trim(); if (!trimmed) continue; if (isSqlEditorTransactionControlStatement(trimmed)) return false; + if (isSqlEditorManagedBlockWrite(type, trimmed)) return false; if (sqlEditorStatementHasManagedWrite(trimmed)) return false; const keyword = resolveSqlEditorOperationKeyword(trimmed); if (!SQL_EDITOR_READ_KEYWORDS.has(keyword)) return false; diff --git a/internal/app/methods_db_multi_test.go b/internal/app/methods_db_multi_test.go index ba64f676..038ed754 100644 --- a/internal/app/methods_db_multi_test.go +++ b/internal/app/methods_db_multi_test.go @@ -21,6 +21,7 @@ type fakeBatchWriteDB struct { lastQuery string lastCtx context.Context queryCalls int + queryQueries []string queryMap map[string][]map[string]interface{} fieldMap map[string][]string messageMap map[string][]string @@ -212,6 +213,7 @@ func (f *fakeBatchWriteDB) ExecContext(ctx context.Context, query string) (int64 func (f *fakeBatchWriteDB) QueryContext(ctx context.Context, query string) ([]map[string]interface{}, []string, error) { f.lastCtx = ctx f.queryCalls++ + f.queryQueries = append(f.queryQueries, query) if err := f.queryErr[query]; err != nil { return nil, nil, err } @@ -1075,6 +1077,58 @@ func TestDBQueryMultiTransactionalKeepsDMLTransactionOpenUntilCommit(t *testing. } } +func TestDBQueryMultiTransactionalKeepsSQLServerBeginEndBlockOpenUntilRollback(t *testing.T) { + originalNewDatabaseFunc := newDatabaseFunc + t.Cleanup(func() { + newDatabaseFunc = originalNewDatabaseFunc + }) + + block := `BEGIN + UPDATE users SET name = 'new' WHERE id = 1; +END;` + fakeDB := &fakeBatchWriteDB{ + execAffected: map[string]int64{block: 1}, + } + newDatabaseFunc = func(dbType string) (db.Database, error) { + return fakeDB, nil + } + + app := NewAppWithSecretStore(secretstore.NewUnavailableStore("test")) + config := connection.ConnectionConfig{Type: "sqlserver", Host: "127.0.0.1", Port: 1433, User: "sa"} + + result := app.DBQueryMultiTransactional(config, "testdb", block, "sqlserver-begin-end-tx-query") + if !result.Success { + t.Fatalf("expected SQL Server BEGIN...END transaction success, got failure: %s", result.Message) + } + if result.TransactionID == "" || !result.TransactionPending { + t.Fatalf("expected pending transaction metadata, got id=%q pending=%v", result.TransactionID, result.TransactionPending) + } + if fakeDB.session == nil { + t.Fatal("expected SQL Server transactional block to open a pinned session") + } + if fakeDB.session.closed { + t.Fatal("expected SQL Server transaction session to stay open before rollback") + } + if !reflect.DeepEqual(fakeDB.execQueries, []string{"BEGIN TRANSACTION"}) { + t.Fatalf("expected SQL Server transaction begin before rollback, got %#v", fakeDB.execQueries) + } + if !reflect.DeepEqual(fakeDB.queryQueries, []string{block}) { + t.Fatalf("expected SQL Server block to execute through the pinned query session, got %#v", fakeDB.queryQueries) + } + + rollbackResult := app.DBRollbackTransaction(result.TransactionID) + if !rollbackResult.Success { + t.Fatalf("expected SQL Server rollback success, got failure: %s", rollbackResult.Message) + } + if !fakeDB.session.closed { + t.Fatal("expected SQL Server transaction session to close after rollback") + } + wantExecs := []string{"BEGIN TRANSACTION", "ROLLBACK TRANSACTION"} + if !reflect.DeepEqual(fakeDB.execQueries, wantExecs) { + t.Fatalf("expected SQL Server rollback without commit, got %#v", fakeDB.execQueries) + } +} + func TestDBQueryMultiTransactionalKeepsTrailingCommentInsideManagedTransaction(t *testing.T) { originalNewDatabaseFunc := newDatabaseFunc t.Cleanup(func() { diff --git a/internal/app/methods_db_transaction.go b/internal/app/methods_db_transaction.go index feec402f..0ae36976 100644 --- a/internal/app/methods_db_transaction.go +++ b/internal/app/methods_db_transaction.go @@ -636,7 +636,7 @@ func shouldUseManagedSQLTransaction(dbType string, query string) bool { if isReadOnlySQLQuery(dbType, stmt) { continue } - if isOracleLikeAnonymousBlockManagedWrite(dbType, stmt) { + if isManagedSQLBlockWrite(dbType, stmt) { hasManagedWrite = true continue } @@ -699,22 +699,27 @@ func isBeginTransactionControlStatement(stmt string, keywordEnd int) bool { } } -func isOracleLikeAnonymousBlockManagedWrite(dbType string, stmt string) bool { - if !isOracleLikeDBType(dbType) { - return false - } - - switch nextSQLSignificantToken(strings.TrimSpace(stmt), 0) { - case "begin", "declare": - return sqlContainsKeyword(stmt, "insert") || - sqlContainsKeyword(stmt, "update") || - sqlContainsKeyword(stmt, "delete") || - sqlContainsKeyword(stmt, "merge") || - sqlContainsKeyword(stmt, "replace") || - sqlContainsKeyword(stmt, "upsert") +func isManagedSQLBlockWrite(dbType string, stmt string) bool { + keyword, keywordEnd := nextSQLKeyword(stmt, 0) + switch { + case isOracleLikeDBType(dbType): + if keyword != "begin" && keyword != "declare" { + return false + } + case isSQLServerDBType(dbType): + if keyword != "begin" || isBeginTransactionControlStatement(stmt, keywordEnd) { + return false + } default: return false } + + return sqlContainsKeyword(stmt, "insert") || + sqlContainsKeyword(stmt, "update") || + sqlContainsKeyword(stmt, "delete") || + sqlContainsKeyword(stmt, "merge") || + sqlContainsKeyword(stmt, "replace") || + sqlContainsKeyword(stmt, "upsert") } func (a *App) DBCommitTransaction(transactionID string) connection.QueryResult { diff --git a/internal/app/methods_db_transaction_test.go b/internal/app/methods_db_transaction_test.go index 6d6e4b53..57c7d41f 100644 --- a/internal/app/methods_db_transaction_test.go +++ b/internal/app/methods_db_transaction_test.go @@ -51,6 +51,46 @@ END;` } } +func TestShouldUseManagedSQLTransaction_SQLServerBeginEndWithDMLUsesManagedTransaction(t *testing.T) { + t.Parallel() + + query := `BEGIN + PRINT 'the word DELETE here is only text'; + -- INSERT INTO audit_logs(id) VALUES (1); + UPDATE users SET name = 'new' WHERE id = 1; +END;` + if !shouldUseManagedSQLTransaction("sqlserver", query) { + t.Fatal("expected SQL Server BEGIN...END block with DML to use SQL editor managed transaction") + } +} + +func TestShouldUseManagedSQLTransaction_SQLServerBeginEndWithoutExecutableDMLStaysUnmanaged(t *testing.T) { + t.Parallel() + + query := `BEGIN + PRINT 'UPDATE users SET name = ''new'''; + -- DELETE FROM users WHERE id = 1; + /* INSERT INTO audit_logs(id) VALUES (1); */ +END;` + if shouldUseManagedSQLTransaction("sqlserver", query) { + t.Fatal("expected SQL Server BEGIN...END block with DML keywords only in comments or strings to stay unmanaged") + } +} + +func TestShouldUseManagedSQLTransaction_SQLServerExplicitBeginStaysUnmanaged(t *testing.T) { + t.Parallel() + + for _, query := range []string{ + "BEGIN TRANSACTION; UPDATE users SET name = 'new' WHERE id = 1; COMMIT TRANSACTION;", + "BEGIN TRAN; UPDATE users SET name = 'new' WHERE id = 1; COMMIT TRAN;", + "BEGIN WORK; UPDATE users SET name = 'new' WHERE id = 1; COMMIT WORK;", + } { + if shouldUseManagedSQLTransaction("sqlserver", query) { + t.Fatalf("expected explicit SQL Server transaction to stay unmanaged: %q", query) + } + } +} + func TestShouldUseManagedSQLTransaction_UsesDialectCommentRules(t *testing.T) { t.Parallel()