🐛 fix(transaction): 修复匿名SQL块自动提交

- 区分匿名 BEGIN...END 块与显式事务控制语句
- 统一前后端方言判定并挂起 Oracle、达梦和 SQL Server 块内 DML
- 补充编辑器、事务控制器及 SQL Server 回滚链路回归测试
This commit is contained in:
Syngnat
2026-07-14 12:22:07 +08:00
parent 72423e80a8
commit 1b3137d947
8 changed files with 254 additions and 35 deletions

View File

@@ -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() {

View File

@@ -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 {

View File

@@ -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()