mirror of
https://github.com/Syngnat/GoNavi.git
synced 2026-08-22 08:53:46 +08:00
🐛 fix(transaction): 修复匿名SQL块自动提交
- 区分匿名 BEGIN...END 块与显式事务控制语句 - 统一前后端方言判定并挂起 Oracle、达梦和 SQL Server 块内 DML - 补充编辑器、事务控制器及 SQL Server 回滚链路回归测试
This commit is contained in:
@@ -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() {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user