This is an automated email from the ASF dual-hosted git repository.

thunguo pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/incubator-seata-go.git


The following commit(s) were added to refs/heads/master by this push:
     new d75876ea feat(sql):generic single-transaction batch helper (#1139)
d75876ea is described below

commit d75876ea38111c3d8a3c8f9e22c5a383c4788f60
Author: Mochimia <[email protected]>
AuthorDate: Sun Aug 23 00:08:14 2026 +0800

    feat(sql):generic single-transaction batch helper (#1139)
---
 pkg/datasource/sql/batch.go               | 205 +++++++++++++++
 pkg/datasource/sql/batch_result.go        | 162 ++++++++++++
 pkg/datasource/sql/batch_seata_at_test.go | 401 ++++++++++++++++++++++++++++++
 pkg/datasource/sql/batch_test.go          | 329 ++++++++++++++++++++++++
 pkg/datasource/sql/conn.go                |  13 +
 pkg/datasource/sql/conn_at.go             |  14 ++
 pkg/datasource/sql/conn_at_test.go        | 206 +++++++++++++++
 pkg/datasource/sql/tx_at.go               |  97 ++++++--
 8 files changed, 1410 insertions(+), 17 deletions(-)

diff --git a/pkg/datasource/sql/batch.go b/pkg/datasource/sql/batch.go
new file mode 100644
index 00000000..a24f5c7c
--- /dev/null
+++ b/pkg/datasource/sql/batch.go
@@ -0,0 +1,205 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package sql
+
+import (
+       "context"
+       gosql "database/sql"
+       "errors"
+       "fmt"
+       "strings"
+)
+
+// The batch helpers in this file are driver-agnostic. They provide ordered
+// execution within one database/sql transaction only. When used with a Seata
+// driver, AT-specific behavior is provided by that driver and its executors.
+
+var (
+       errNilBatchDB            = errors.New("batch db is nil")
+       errNilBatchTx            = errors.New("batch tx is nil")
+       errEmptyBatchQuery       = errors.New("batch query is empty")
+       errInconsistentBatchArgs = errors.New("inconsistent batch argument 
count")
+)
+
+// batchExecContext describes one semantic batch.
+// A batch contains exactly one SQL template and an ordered set of argument 
groups.
+// Its lifetime is limited to one batch invocation.
+type batchExecContext struct {
+       query     string
+       batchArgs [][]any
+}
+
+func newBatchExecContext(ctx context.Context, query string, batchArgs [][]any) 
(*batchExecContext, error) {
+       if err := ctx.Err(); err != nil {
+               return nil, err
+       }
+
+       if strings.TrimSpace(query) == "" {
+               return nil, errEmptyBatchQuery
+       }
+
+       if len(batchArgs) > 1 {
+               expectedArgCount := len(batchArgs[0])
+
+               for i := 1; i < len(batchArgs); i++ {
+                       if len(batchArgs[i]) != expectedArgCount {
+                               return nil, fmt.Errorf(
+                                       "%w: batch item %d has %d arguments, 
expected %d",
+                                       errInconsistentBatchArgs, i, 
len(batchArgs[i]), expectedArgCount,
+                               )
+                       }
+               }
+       }
+
+       return &batchExecContext{query: query, batchArgs: batchArgs}, nil
+}
+
+// ExecBatchContext executes one SQL template with multiple argument groups.
+//
+// The transaction is owned by this function.
+// All batch items are executed sequentially in one local transaction.
+// The first execution error stops the batch and causes the whole transaction 
to be rolled back.
+// The returned result remains valid on error and contains one ordered item 
for each argument group.
+//
+// When used with a Seata AT driver in a global transaction, all batch items
+// participate in the same local transaction and therefore share the same AT
+// branch lifecycle. AT-specific image and undo-log handling remains the
+// responsibility of the Seata driver and its executors.
+//
+// This is the batch counterpart of database/sql.DB.ExecContext:
+// when callers need to combine the batch with other statements in the same 
transaction,
+// they should begin a transaction explicitly and use ExecBatchInTxContext.
+func ExecBatchContext(ctx context.Context, db *gosql.DB, query string, 
batchArgs [][]any) (BatchResult, error) {
+       result := newBatchResult(len(batchArgs), BatchTransactionNotStarted)
+       if db == nil {
+               result.Outcome.FailurePhase = BatchPhaseValidate
+               return result, newBatchError(result, errNilBatchDB, nil)
+       }
+
+       batchCtx, err := newBatchExecContext(ctx, query, batchArgs)
+       if err != nil {
+               result.Outcome.FailurePhase = BatchPhaseValidate
+               return result, newBatchError(result, err, nil)
+       }
+
+       if len(batchCtx.batchArgs) == 0 {
+               return result, nil
+       }
+
+       tx, err := db.BeginTx(ctx, nil)
+       if err != nil {
+               result.Outcome.FailurePhase = BatchPhaseBegin
+               cause := fmt.Errorf("begin batch transaction: %w", err)
+               return result, newBatchError(result, cause, nil)
+       }
+       defer func() {
+               _ = tx.Rollback()
+       }()
+
+       // database/sql may already have rolled back the transaction after 
context cancellation.
+       // Do not let ErrTxDone hide the original execution error.
+       if failedIndex, err := executeBatch(ctx, tx, batchCtx, &result); err != 
nil {
+               result.Outcome.FailedIndex = failedIndex
+               result.Outcome.FailurePhase = BatchPhaseExecute
+               rollbackErr := tx.Rollback()
+               if rollbackErr != nil && !errors.Is(rollbackErr, 
gosql.ErrTxDone) {
+                       result.Outcome.TransactionState = 
BatchTransactionRollbackFailed
+                       rollbackErr = fmt.Errorf("rollback batch transaction: 
%w", rollbackErr)
+                       return result, newBatchError(result, err, rollbackErr)
+               }
+               result.Outcome.TransactionState = BatchTransactionRolledBack
+               return result, newBatchError(result, err, nil)
+       }
+
+       if err := tx.Commit(); err != nil {
+               result.Outcome.FailurePhase = BatchPhaseCommit
+               result.Outcome.TransactionState = BatchTransactionCommitUnknown
+               var rollbackErr error
+               var commitErr *atCommitError
+               if errors.As(err, &commitErr) {
+                       switch commitErr.outcome {
+                       case atCommitOutcomeRolledBack:
+                               result.Outcome.TransactionState = 
BatchTransactionRolledBack
+                       case atCommitOutcomeRollbackFailed:
+                               result.Outcome.TransactionState = 
BatchTransactionRollbackFailed
+                               rollbackErr = commitErr.rollbackErr
+                       case atCommitOutcomeCommitted:
+                               result.Outcome.TransactionState = 
BatchTransactionCommitted
+                       }
+               }
+               cause := fmt.Errorf("commit batch transaction: %w", err)
+               return result, newBatchError(result, cause, rollbackErr)
+       }
+       result.Outcome.TransactionState = BatchTransactionCommitted
+       return result, nil
+}
+
+// ExecBatchInTxContext executes one SQL template with multiple argument groups
+// inside a caller-owned transaction.
+//
+// When used with a Seata AT driver, the batch joins the caller's existing 
transaction
+// and doesn't create or finish an AT branch on its own.
+//
+// The function never commits or rolls back tx.
+// If an item fails, execution stops immediately and the error is returned to 
the caller,
+// which remains responsible for the transaction lifecycle.
+// The returned transaction state is pending because the caller owns its final 
outcome.
+func ExecBatchInTxContext(ctx context.Context, tx *gosql.Tx, query string, 
batchArgs [][]any) (BatchResult, error) {
+       result := newBatchResult(len(batchArgs), BatchTransactionNotStarted)
+       if tx == nil {
+               result.Outcome.FailurePhase = BatchPhaseValidate
+               return result, newBatchError(result, errNilBatchTx, nil)
+       }
+       result.Outcome.TransactionState = BatchTransactionPending
+
+       batchCtx, err := newBatchExecContext(ctx, query, batchArgs)
+       if err != nil {
+               result.Outcome.FailurePhase = BatchPhaseValidate
+               return result, newBatchError(result, err, nil)
+       }
+
+       if len(batchCtx.batchArgs) == 0 {
+               return result, nil
+       }
+
+       failedIndex, err := executeBatch(ctx, tx, batchCtx, &result)
+       if err != nil {
+               result.Outcome.FailedIndex = failedIndex
+               result.Outcome.FailurePhase = BatchPhaseExecute
+               return result, newBatchError(result, err, nil)
+       }
+       return result, nil
+}
+
+// executeBatch is the semantic batch execution core.
+//
+// Regardless of whether the transaction was created by ExecBatchContext or
+// supplied by the caller, all items are executed on the same *sql.Tx and
+// therefore the same underlying database transaction.
+func executeBatch(ctx context.Context, tx *gosql.Tx, batchCtx 
*batchExecContext, result *BatchResult) (int, error) {
+       for i, arg := range batchCtx.batchArgs {
+               sqlResult, err := tx.ExecContext(ctx, batchCtx.query, arg...)
+               if err != nil {
+                       result.Items[i].State = BatchItemFailed
+                       result.Items[i].execErr = err
+                       return i, fmt.Errorf("execute batch item %d: %w", i, 
err)
+               }
+               result.Items[i].recordResult(sqlResult)
+       }
+       return NoFailedBatchItem, nil
+}
diff --git a/pkg/datasource/sql/batch_result.go 
b/pkg/datasource/sql/batch_result.go
new file mode 100644
index 00000000..aff37440
--- /dev/null
+++ b/pkg/datasource/sql/batch_result.go
@@ -0,0 +1,162 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package sql
+
+import (
+       gosql "database/sql"
+       "errors"
+       "fmt"
+)
+
+// NoFailedBatchItem indicates that a batch failure is not associated with an 
item.
+const NoFailedBatchItem = -1
+
+// BatchPhase identifies the phase that caused a batch to fail.
+type BatchPhase uint8
+
+const (
+       BatchPhaseNone BatchPhase = iota
+       BatchPhaseValidate
+       BatchPhaseBegin
+       BatchPhaseExecute
+       BatchPhaseCommit
+)
+
+// BatchTransactionState describes the transaction outcome known to the batch 
API.
+type BatchTransactionState uint8
+
+const (
+       BatchTransactionNotStarted BatchTransactionState = iota
+       BatchTransactionPending
+       BatchTransactionCommitted
+       BatchTransactionRolledBack
+       BatchTransactionRollbackFailed
+       BatchTransactionCommitUnknown
+)
+
+// BatchItemState describes the execution state of one ordered batch item.
+type BatchItemState uint8
+
+const (
+       BatchItemNotExecuted BatchItemState = iota
+       BatchItemExecuted
+       BatchItemFailed
+)
+
+// BatchOutcome describes where a batch failed and its transaction outcome.
+type BatchOutcome struct {
+       FailedIndex      int
+       FailurePhase     BatchPhase
+       TransactionState BatchTransactionState
+}
+
+// BatchResult contains one result for each input argument group, in input 
order.
+type BatchResult struct {
+       Items   []ItemResult
+       Outcome BatchOutcome
+}
+
+// ItemResult contains the execution result of one batch item.
+// BatchItemExecuted means execution succeeded; durability is described by 
BatchResult.Outcome.
+type ItemResult struct {
+       Index int
+       State BatchItemState
+
+       execErr         error
+       lastInsertID    int64
+       lastInsertIDErr error
+       rowsAffected    int64
+       rowsAffectedErr error
+}
+
+var errBatchItemResultUnavailable = errors.New("batch item result is 
unavailable")
+
+// Err returns the item's execution error, if any.
+func (r ItemResult) Err() error {
+       return r.execErr
+}
+
+// LastInsertId returns the driver's snapshotted LastInsertId result.
+func (r ItemResult) LastInsertId() (int64, error) {
+       if r.State != BatchItemExecuted {
+               return 0, errBatchItemResultUnavailable
+       }
+       return r.lastInsertID, r.lastInsertIDErr
+}
+
+// RowsAffected returns the driver's snapshotted RowsAffected result.
+func (r ItemResult) RowsAffected() (int64, error) {
+       if r.State != BatchItemExecuted {
+               return 0, errBatchItemResultUnavailable
+       }
+       return r.rowsAffected, r.rowsAffectedErr
+}
+
+func (r *ItemResult) recordResult(result gosql.Result) {
+       r.State = BatchItemExecuted
+       r.lastInsertID, r.lastInsertIDErr = result.LastInsertId()
+       r.rowsAffected, r.rowsAffectedErr = result.RowsAffected()
+}
+
+// BatchError describes a batch failure while preserving its underlying errors.
+type BatchError struct {
+       Outcome     BatchOutcome
+       Cause       error
+       RollbackErr error
+}
+
+func (e *BatchError) Error() string {
+       message := "batch failed"
+       if e.Cause != nil {
+               message = e.Cause.Error()
+       }
+       if e.RollbackErr != nil {
+               return fmt.Sprintf("%s; %v", message, e.RollbackErr)
+       }
+       return message
+}
+
+// Unwrap exposes both the primary failure and a rollback failure to errors.Is 
and errors.As.
+func (e *BatchError) Unwrap() []error {
+       errs := make([]error, 0, 2)
+       if e.Cause != nil {
+               errs = append(errs, e.Cause)
+       }
+       if e.RollbackErr != nil {
+               errs = append(errs, e.RollbackErr)
+       }
+       return errs
+}
+
+func newBatchResult(itemCount int, state BatchTransactionState) BatchResult {
+       result := BatchResult{
+               Items: make([]ItemResult, itemCount),
+               Outcome: BatchOutcome{
+                       FailedIndex:      NoFailedBatchItem,
+                       TransactionState: state,
+               },
+       }
+       for i := range result.Items {
+               result.Items[i].Index = i
+       }
+       return result
+}
+
+func newBatchError(result BatchResult, cause, rollbackErr error) *BatchError {
+       return &BatchError{Outcome: result.Outcome, Cause: cause, RollbackErr: 
rollbackErr}
+}
diff --git a/pkg/datasource/sql/batch_seata_at_test.go 
b/pkg/datasource/sql/batch_seata_at_test.go
new file mode 100644
index 00000000..e0755a48
--- /dev/null
+++ b/pkg/datasource/sql/batch_seata_at_test.go
@@ -0,0 +1,401 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package sql
+
+import (
+       "context"
+       gosql "database/sql"
+       "database/sql/driver"
+       "errors"
+       "io"
+       "strings"
+       "sync/atomic"
+       "testing"
+
+       "github.com/golang/mock/gomock"
+       "github.com/google/uuid"
+       "github.com/stretchr/testify/require"
+
+       "seata.apache.org/seata-go/v2/pkg/datasource/sql/datasource"
+       "seata.apache.org/seata-go/v2/pkg/datasource/sql/mock"
+       "seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+       "seata.apache.org/seata-go/v2/pkg/datasource/sql/undo"
+       undoparser "seata.apache.org/seata-go/v2/pkg/datasource/sql/undo/parser"
+       "seata.apache.org/seata-go/v2/pkg/protocol/branch"
+       "seata.apache.org/seata-go/v2/pkg/rm"
+       "seata.apache.org/seata-go/v2/pkg/tm"
+)
+
+type batchATRows struct {
+       columns []string
+       data    [][]driver.Value
+       index   int
+}
+
+func (r *batchATRows) Columns() []string { return r.columns }
+func (r *batchATRows) Close() error      { return nil }
+func (r *batchATRows) Next(dest []driver.Value) error {
+       if r.index == len(r.data) {
+               return io.EOF
+       }
+       copy(dest, r.data[r.index])
+       r.index++
+       return nil
+}
+
+type batchATTableCache struct{}
+
+func (batchATTableCache) Init(context.Context, *gosql.DB) error { return nil }
+func (batchATTableCache) Destroy() error                        { return nil }
+func (batchATTableCache) GetTableMeta(_ context.Context, _, table string) 
(*types.TableMeta, error) {
+       idColumn := types.ColumnMeta{ColumnName: "id", DatabaseTypeString: 
"BIGINT"}
+       return &types.TableMeta{
+               TableName:   table,
+               ColumnNames: []string{"id", "balance"},
+               Columns: map[string]types.ColumnMeta{
+                       "id":      idColumn,
+                       "balance": {ColumnName: "balance", DatabaseTypeString: 
"BIGINT"},
+               },
+               Indexs: map[string]types.IndexMeta{
+                       "PRIMARY": {IType: types.IndexTypePrimaryKey, Columns: 
[]types.ColumnMeta{idColumn}},
+               },
+       }, nil
+}
+
+func newBatchSeataATTestDB(t *testing.T,
+       ctrl *gomock.Controller,
+) (*gosql.DB, *mock.MockTestDriverConn, *mock.MockTestDriverTx, 
*mock.MockDataSourceManager) {
+       t.Helper()
+
+       mockMgr := mock.NewMockDataSourceManager(ctrl)
+       mockMgr.SetBranchType(branch.BranchTypeAT)
+       mockMgr.EXPECT().RegisterResource(gomock.Any()).Times(1).Return(nil)
+       registerResourceManagerForTest(t, mockMgr)
+
+       mockTx := mock.NewMockTestDriverTx(ctrl)
+       mockConn := mock.NewMockTestDriverConn(ctrl)
+
+       mockConn.EXPECT().
+               QueryContext(gomock.Any(), "SELECT VERSION()", gomock.Any()).
+               AnyTimes().
+               DoAndReturn(func(ctx context.Context, query string, args 
[]driver.NamedValue) (driver.Rows, error) {
+                       rows := &mysqlMockRows{}
+                       rows.data = [][]interface{}{
+                               {"8.0.29"},
+                       }
+                       return rows, nil
+               })
+
+       mockConn.EXPECT().ResetSession(gomock.Any()).AnyTimes().Return(nil)
+       mockConn.EXPECT().Close().AnyTimes().Return(nil)
+
+       connector := mock.NewMockTestDriverConnector(ctrl)
+       connector.EXPECT().Connect(gomock.Any()).AnyTimes().Return(mockConn, 
nil)
+
+       targetDB := gosql.OpenDB(connector)
+       t.Cleanup(func() {
+               _ = targetDB.Close()
+       })
+
+       previousTableCache := datasource.GetTableCache(types.DBTypeMySQL)
+       t.Cleanup(func() {
+               datasource.RegisterTableCache(types.DBTypeMySQL, 
previousTableCache)
+       })
+
+       proxyConnector, err := (&seataDriver{
+               branchType: branch.BranchTypeAT,
+               transType:  types.ATMode,
+               descriptor: mySQLDriverDescriptor,
+               target:     mySQLDriverDescriptor.target,
+               targetName: "mysql",
+       }).getOpenConnectorProxy(
+               connector,
+               types.DBTypeMySQL,
+               targetDB,
+               
"root:password@tcp(mock:3306)/seata_client?multiStatements=true",
+       )
+       require.NoError(t, err)
+
+       baseConnector, ok := proxyConnector.(*seataConnector)
+       require.True(t, ok)
+       db := gosql.OpenDB(&seataATConnector{seataConnector: baseConnector})
+
+       return db, mockConn, mockTx, mockMgr
+}
+
+func expectBatchATImageQueries(t *testing.T, mockConn 
*mock.MockTestDriverConn, snapshots [][]driver.Value) {
+       t.Helper()
+
+       callIndex := 0
+       mockConn.EXPECT().QueryContext(gomock.Any(), gomock.Not("SELECT 
VERSION()"), gomock.Any()).
+               Times(len(snapshots)).DoAndReturn(func(_ context.Context, query 
string, args []driver.NamedValue) (driver.Rows, error) {
+               require.Len(t, args, 1)
+               require.Equal(t, snapshots[callIndex][0], args[0].Value)
+               if callIndex%2 == 0 {
+                       require.Contains(t, query, "FOR UPDATE")
+               } else {
+                       require.NotContains(t, query, "FOR UPDATE")
+               }
+               rows := &batchATRows{columns: []string{"id", "balance"}, data: 
[][]driver.Value{snapshots[callIndex]}}
+               callIndex++
+               return rows, nil
+       })
+}
+
+func TestExecBatchContextWithSeataATDriverUsesSingleBranchLifecycle(t 
*testing.T) {
+       CleanTxHooks()
+       t.Cleanup(CleanTxHooks)
+
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       db, mockConn, mockTx, mockMgr := newBatchSeataATTestDB(t, ctrl)
+       defer db.Close()
+       datasource.RegisterTableCache(types.DBTypeMySQL, batchATTableCache{})
+
+       previousUndoConfig := undo.UndoConfig
+       undo.UndoConfig = undo.Config{LogSerialization: "json", LogTable: 
"undo_log"}
+       t.Cleanup(func() { undo.UndoConfig = previousUndoConfig })
+
+       ctx := tm.InitSeataContext(context.Background())
+       xid := uuid.NewString()
+       tm.SetXID(ctx, xid)
+
+       query := "UPDATE account SET balance = ? WHERE id = ?"
+       var txCtx *types.TransactionContext
+       RegisterTxHook(&mockTxHook{beforeCommit: func(tx *Tx) error {
+               txCtx = tx.tranCtx
+               return nil
+       }})
+
+       mockConn.EXPECT().BeginTx(gomock.Any(), 
gomock.Any()).Times(1).Return(mockTx, nil)
+       expectBatchATImageQueries(t, mockConn, [][]driver.Value{
+               {int64(1), int64(100)}, {int64(1), int64(110)},
+               {int64(2), int64(200)}, {int64(2), int64(220)},
+       })
+       mockConn.EXPECT().ExecContext(gomock.Any(), query, 
gomock.Any()).Times(2).DoAndReturn(
+               func(_ context.Context, _ string, args []driver.NamedValue) 
(driver.Result, error) {
+                       require.Len(t, args, 2)
+                       return driver.RowsAffected(1), nil
+               },
+       )
+
+       const branchID = int64(123)
+       var registeredLockKeys []string
+       registerCall := mockMgr.EXPECT().BranchRegister(gomock.Any(), 
gomock.Any()).Times(1).DoAndReturn(
+               func(_ context.Context, param rm.BranchRegisterParam) (int64, 
error) {
+                       require.Equal(t, xid, param.Xid)
+                       registeredLockKeys = strings.FieldsFunc(param.LockKeys, 
func(r rune) bool { return r == ';' })
+                       return branchID, nil
+               },
+       )
+
+       undoStmt := mock.NewMockTestDriverStmt(ctrl)
+       prepareUndoCall := mockConn.EXPECT().PrepareContext(gomock.Any(), 
gomock.Any()).Times(1).DoAndReturn(
+               func(_ context.Context, query string) (driver.Stmt, error) {
+                       require.Contains(t, query, "INSERT INTO undo_log")
+                       return undoStmt, nil
+               },
+       )
+       undoStmt.EXPECT().Close().Times(1).Return(nil)
+       var branchUndoLog *undo.BranchUndoLog
+       flushUndoCall := undoStmt.EXPECT().ExecContext(gomock.Any(), 
gomock.Any()).Times(1).DoAndReturn(
+               func(_ context.Context, args []driver.NamedValue) 
(driver.Result, error) {
+                       require.Len(t, args, 5)
+                       require.EqualValues(t, branchID, args[0].Value)
+                       require.Equal(t, xid, args[1].Value)
+                       rollbackInfo, ok := args[3].Value.([]byte)
+                       require.True(t, ok)
+                       var err error
+                       branchUndoLog, err = 
(&undoparser.JsonParser{}).Decode(rollbackInfo)
+                       require.NoError(t, err)
+                       return driver.ResultNoRows, nil
+               },
+       )
+
+       commitCall := mockTx.EXPECT().Commit().Times(1).Return(nil)
+       reportCall := mockMgr.EXPECT().BranchReport(gomock.Any(), 
gomock.Any()).Times(1).DoAndReturn(
+               func(_ context.Context, param rm.BranchReportParam) error {
+                       require.EqualValues(t, branchID, param.BranchId)
+                       require.EqualValues(t, branch.BranchStatusPhaseoneDone, 
param.Status)
+                       return nil
+               },
+       )
+       gomock.InOrder(registerCall, prepareUndoCall, flushUndoCall, 
commitCall, reportCall)
+
+       result, err := ExecBatchContext(ctx, db, query, [][]any{{int64(110), 
int64(1)}, {int64(220), int64(2)}})
+
+       require.NoError(t, err)
+       require.Equal(t, BatchTransactionCommitted, 
result.Outcome.TransactionState)
+       require.NotNil(t, txCtx)
+       require.Len(t, txCtx.RoundImages.BeofreImages(), 2)
+       require.Len(t, txCtx.RoundImages.AfterImages(), 2)
+
+       lockKeys := make([]string, 0, len(txCtx.LockKeys))
+       for lockKey := range txCtx.LockKeys {
+               lockKeys = append(lockKeys, lockKey)
+       }
+       require.ElementsMatch(t, []string{"ACCOUNT:1", "ACCOUNT:2"}, lockKeys)
+       require.ElementsMatch(t, lockKeys, registeredLockKeys)
+
+       require.NotNil(t, branchUndoLog)
+       require.Equal(t, xid, branchUndoLog.Xid)
+       require.EqualValues(t, branchID, branchUndoLog.BranchID)
+       require.Len(t, branchUndoLog.Logs, 2)
+       for i, expectedID := range []int64{1, 2} {
+               require.EqualValues(t, expectedID, 
branchUndoLog.Logs[i].BeforeImage.Rows[0].GetColumnMap()["id"].Value)
+               require.EqualValues(t, expectedID, 
branchUndoLog.Logs[i].AfterImage.Rows[0].GetColumnMap()["id"].Value)
+       }
+}
+
+func TestExecBatchContextWithSeataATDriverRollsBackOwnedTransactionOnFailure(t 
*testing.T) {
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       db, mockConn, mockTx, mockMgr := newBatchSeataATTestDB(t, ctrl)
+       defer db.Close()
+       datasource.RegisterTableCache(types.DBTypeMySQL, batchATTableCache{})
+
+       ctx := tm.InitSeataContext(context.Background())
+       tm.SetXID(ctx, uuid.NewString())
+
+       query := "UPDATE account SET balance = ? WHERE id = ?"
+       execErr := errors.New("execute failed")
+
+       var execCount int32
+       mockConn.EXPECT().BeginTx(gomock.Any(), 
gomock.Any()).Times(1).Return(mockTx, nil)
+       expectBatchATImageQueries(t, mockConn, [][]driver.Value{
+               {int64(1), int64(100)}, {int64(1), int64(110)}, {int64(2), 
int64(200)},
+       })
+       mockConn.EXPECT().ExecContext(gomock.Any(), query, gomock.Any()).
+               Times(2).DoAndReturn(func(_ context.Context, _ string, args 
[]driver.NamedValue) (driver.Result, error) {
+               count := atomic.AddInt32(&execCount, 1)
+               if count == 2 {
+                       return nil, execErr
+               }
+               return driver.RowsAffected(1), nil
+       })
+
+       mockMgr.EXPECT().BranchRegister(gomock.Any(), gomock.Any()).Times(0)
+       mockMgr.EXPECT().BranchReport(gomock.Any(), gomock.Any()).Times(0)
+       mockConn.EXPECT().PrepareContext(gomock.Any(), gomock.Any()).Times(0)
+       mockTx.EXPECT().Commit().Times(0)
+       mockTx.EXPECT().Rollback().Times(1).Return(nil)
+       result, err := ExecBatchContext(ctx, db, query, [][]any{
+               {int64(110), int64(1)}, {int64(220), int64(2)}, {int64(330), 
int64(3)},
+       })
+
+       require.Error(t, err)
+       require.ErrorIs(t, err, execErr)
+       require.Contains(t, err.Error(), "batch item 1")
+       require.Equal(t, BatchItemExecuted, result.Items[0].State)
+       require.Equal(t, BatchItemFailed, result.Items[1].State)
+       require.Equal(t, BatchItemNotExecuted, result.Items[2].State)
+
+       // item2 must never execute
+       require.Equal(t, int32(2), atomic.LoadInt32(&execCount))
+}
+
+func TestExecBatchContextMapsATPreCommitOutcome(t *testing.T) {
+       for _, test := range []struct {
+               name          string
+               rollbackErr   error
+               expectedState BatchTransactionState
+       }{
+               {name: "rolled back", expectedState: 
BatchTransactionRolledBack},
+               {name: "rollback failed", rollbackErr: errors.New("rollback 
failed"), expectedState: BatchTransactionRollbackFailed},
+       } {
+               t.Run(test.name, func(t *testing.T) {
+                       CleanTxHooks()
+                       t.Cleanup(CleanTxHooks)
+
+                       ctrl := gomock.NewController(t)
+                       defer ctrl.Finish()
+
+                       db, mockConn, mockTx, _ := newBatchSeataATTestDB(t, 
ctrl)
+                       defer db.Close()
+
+                       commitErr := errors.New("before commit failed")
+                       RegisterTxHook(&mockTxHook{beforeCommit: func(*Tx) 
error { return commitErr }})
+
+                       mockConn.EXPECT().BeginTx(gomock.Any(), 
gomock.Any()).Times(1).Return(mockTx, nil)
+                       mockConn.EXPECT().ExecContext(gomock.Any(), "SELECT ?", 
gomock.Any()).Times(1).Return(driver.ResultNoRows, nil)
+                       
mockTx.EXPECT().Rollback().Times(1).Return(test.rollbackErr)
+
+                       ctx := tm.InitSeataContext(context.Background())
+                       tm.SetXID(ctx, uuid.NewString())
+                       result, err := ExecBatchContext(ctx, db, "SELECT ?", 
[][]any{{"item0"}})
+
+                       require.ErrorIs(t, err, commitErr)
+                       require.Equal(t, BatchPhaseCommit, 
result.Outcome.FailurePhase)
+                       require.Equal(t, NoFailedBatchItem, 
result.Outcome.FailedIndex)
+                       require.Equal(t, test.expectedState, 
result.Outcome.TransactionState)
+                       require.Equal(t, BatchItemExecuted, 
result.Items[0].State)
+
+                       var batchErr *BatchError
+                       require.ErrorAs(t, err, &batchErr)
+                       if test.rollbackErr == nil {
+                               require.NoError(t, batchErr.RollbackErr)
+                       } else {
+                               require.ErrorIs(t, err, test.rollbackErr)
+                               require.ErrorIs(t, batchErr.RollbackErr, 
test.rollbackErr)
+                       }
+               })
+       }
+}
+
+func 
TestExecBatchInTxContextWithSeataATDriverAllowsFollowingExecInSameTransaction(t 
*testing.T) {
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       db, mockConn, mockTx, _ := newBatchSeataATTestDB(t, ctrl)
+       defer db.Close()
+
+       ctx := tm.InitSeataContext(context.Background())
+       tm.SetXID(ctx, uuid.NewString())
+
+       batchQuery := "SELECT ?"
+       normalQuery := "SELECT ?"
+
+       var executedArgs []any
+
+       // Caller creates one transaction.
+       mockConn.EXPECT().BeginTx(gomock.Any(), 
gomock.Any()).Times(1).Return(mockTx, nil)
+
+       // 2 batch items + 1 ordinary SQL .
+       mockConn.EXPECT().ExecContext(gomock.Any(), gomock.Any(), gomock.Any()).
+               Times(3).
+               DoAndReturn(func(ctx context.Context, query string, args 
[]driver.NamedValue) (driver.Result, error) {
+                       executedArgs = append(executedArgs, args[0].Value)
+                       return driver.ResultNoRows, nil
+               })
+
+       mockTx.EXPECT().Commit().Times(1).Return(nil)
+
+       tx, err := db.BeginTx(ctx, nil)
+       require.NoError(t, err)
+       _, err = ExecBatchInTxContext(ctx, tx, batchQuery, [][]any{{"item0"}, 
{"item1"}})
+       require.NoError(t, err)
+
+       // Batch execution must not close caller-owned transaction
+       _, err = tx.ExecContext(ctx, normalQuery, "normal")
+       require.NoError(t, err)
+
+       require.NoError(t, tx.Commit())
+       require.Equal(t, []any{"item0", "item1", "normal"}, executedArgs)
+}
diff --git a/pkg/datasource/sql/batch_test.go b/pkg/datasource/sql/batch_test.go
new file mode 100644
index 00000000..c6de089b
--- /dev/null
+++ b/pkg/datasource/sql/batch_test.go
@@ -0,0 +1,329 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package sql
+
+import (
+       "context"
+       "database/sql/driver"
+       "errors"
+       "regexp"
+       "testing"
+
+       "github.com/DATA-DOG/go-sqlmock"
+       "github.com/stretchr/testify/require"
+)
+
+type panickingBatchValuer struct{}
+
+func (panickingBatchValuer) Value() (driver.Value, error) {
+       panic("batch valuer panic")
+}
+
+func TestExecBatchContextRejectsInconsistentArgumentCount(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       ctx := context.Background()
+       query := "UPDATE user SET name = ? WHERE id = ?"
+
+       result, err := ExecBatchContext(ctx, db, query, [][]any{{"user1", 1}, 
{"user2"}, {"user3", 3}})
+
+       require.ErrorIs(t, err, errInconsistentBatchArgs)
+       require.Contains(t, err.Error(), "batch item 1")
+       require.Contains(t, err.Error(), "has 1 arguments, expected 2")
+       require.Equal(t, BatchPhaseValidate, result.Outcome.FailurePhase)
+       require.Equal(t, BatchTransactionNotStarted, 
result.Outcome.TransactionState)
+
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchContextCommitOnSuccess(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       ctx := context.Background()
+       query := "UPDATE user SET name = ? WHERE id = ?"
+       metadataErr := errors.New("result metadata unavailable")
+
+       mock.ExpectBegin()
+
+       mock.ExpectExec(regexp.QuoteMeta(query)).
+               WithArgs("user1", 1).
+               WillReturnResult(sqlmock.NewResult(10, 1))
+
+       mock.ExpectExec(regexp.QuoteMeta(query)).
+               WithArgs("user2", 2).
+               WillReturnResult(sqlmock.NewResult(20, 2))
+
+       mock.ExpectExec(regexp.QuoteMeta(query)).
+               WithArgs("user3", 3).
+               WillReturnResult(sqlmock.NewErrorResult(metadataErr))
+
+       mock.ExpectCommit()
+
+       result, err := ExecBatchContext(ctx, db, query, [][]any{{"user1", 1}, 
{"user2", 2}, {"user3", 3}})
+       require.NoError(t, err)
+       require.Equal(t, BatchTransactionCommitted, 
result.Outcome.TransactionState)
+       require.Equal(t, NoFailedBatchItem, result.Outcome.FailedIndex)
+       require.Len(t, result.Items, 3)
+
+       for i := range result.Items {
+               require.Equal(t, i, result.Items[i].Index)
+               require.Equal(t, BatchItemExecuted, result.Items[i].State)
+       }
+
+       lastInsertID, err := result.Items[0].LastInsertId()
+       require.NoError(t, err)
+       require.EqualValues(t, 10, lastInsertID)
+       rowsAffected, err := result.Items[1].RowsAffected()
+       require.NoError(t, err)
+       require.EqualValues(t, 2, rowsAffected)
+       _, err = result.Items[2].LastInsertId()
+       require.ErrorIs(t, err, metadataErr)
+       _, err = result.Items[2].RowsAffected()
+       require.ErrorIs(t, err, metadataErr)
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchContextRollbackOnItemFailure(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       ctx := context.Background()
+       query := "UPDATE user SET name = ? WHERE id = ?"
+       execErr := errors.New("execute failed")
+
+       mock.ExpectBegin()
+
+       mock.ExpectExec(regexp.QuoteMeta(query)).
+               WithArgs("user1", 1).
+               WillReturnResult(sqlmock.NewResult(0, 1))
+
+       mock.ExpectExec(regexp.QuoteMeta(query)).
+               WithArgs("user2", 2).
+               WillReturnError(execErr)
+
+       mock.ExpectRollback()
+
+       result, err := ExecBatchContext(ctx, db, query, [][]any{{"user1", 1}, 
{"user2", 2}, {"user3", 3}})
+
+       require.Error(t, err)
+       require.ErrorIs(t, err, execErr)
+       require.Contains(t, err.Error(), "batch item 1")
+       require.Equal(t, BatchOutcome{
+               FailedIndex:      1,
+               FailurePhase:     BatchPhaseExecute,
+               TransactionState: BatchTransactionRolledBack,
+       }, result.Outcome)
+       require.Equal(t, BatchItemExecuted, result.Items[0].State)
+       require.Equal(t, BatchItemFailed, result.Items[1].State)
+       require.Equal(t, BatchItemNotExecuted, result.Items[2].State)
+       require.ErrorIs(t, result.Items[1].Err(), execErr)
+       rowsAffected, resultErr := result.Items[0].RowsAffected()
+       require.NoError(t, resultErr)
+       require.EqualValues(t, 1, rowsAffected)
+
+       var batchErr *BatchError
+       require.ErrorAs(t, err, &batchErr)
+       require.Equal(t, result.Outcome, batchErr.Outcome)
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchContextReportsRollbackFailure(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       execErr := errors.New("execute failed")
+       rollbackErr := errors.New("rollback failed")
+       mock.ExpectBegin()
+       mock.ExpectExec("UPDATE").WillReturnError(execErr)
+       mock.ExpectRollback().WillReturnError(rollbackErr)
+
+       result, err := ExecBatchContext(context.Background(), db, "UPDATE", 
[][]any{{}})
+
+       require.ErrorIs(t, err, execErr)
+       require.ErrorIs(t, err, rollbackErr)
+       require.Equal(t, BatchTransactionRollbackFailed, 
result.Outcome.TransactionState)
+       require.Equal(t, 0, result.Outcome.FailedIndex)
+
+       var batchErr *BatchError
+       require.ErrorAs(t, err, &batchErr)
+       require.ErrorIs(t, batchErr.RollbackErr, rollbackErr)
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchContextPreservesResultsWhenCommitOutcomeUnknown(t 
*testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       commitErr := errors.New("commit failed")
+       mock.ExpectBegin()
+       mock.ExpectExec("UPDATE").WillReturnResult(sqlmock.NewResult(10, 2))
+       mock.ExpectCommit().WillReturnError(commitErr)
+
+       result, err := ExecBatchContext(context.Background(), db, "UPDATE", 
[][]any{{}})
+
+       require.ErrorIs(t, err, commitErr)
+       require.Equal(t, BatchPhaseCommit, result.Outcome.FailurePhase)
+       require.Equal(t, BatchTransactionCommitUnknown, 
result.Outcome.TransactionState)
+       require.Equal(t, NoFailedBatchItem, result.Outcome.FailedIndex)
+       require.Equal(t, BatchItemExecuted, result.Items[0].State)
+       rowsAffected, resultErr := result.Items[0].RowsAffected()
+       require.NoError(t, resultErr)
+       require.EqualValues(t, 2, rowsAffected)
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchContextRollsBackWhenExecutionPanics(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       mock.ExpectBegin()
+       mock.ExpectRollback()
+
+       require.PanicsWithValue(t, "batch valuer panic", func() {
+               _, _ = ExecBatchContext(context.Background(), db, "UPDATE", 
[][]any{{panickingBatchValuer{}}})
+       })
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchInTxContextKeepsCallerTransactionOpen(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       ctx := context.Background()
+       batchQuery := "UPDATE user SET name = ? WHERE id=?"
+       singleQuery := "UPDATE account SET balance = ? WHERE id = ?"
+
+       mock.ExpectBegin()
+       tx, err := db.BeginTx(ctx, nil)
+       require.NoError(t, err)
+
+       mock.ExpectExec(regexp.QuoteMeta(batchQuery)).
+               WithArgs("user1", 1).
+               WillReturnResult(sqlmock.NewResult(0, 1))
+
+       mock.ExpectExec(regexp.QuoteMeta(batchQuery)).
+               WithArgs("user2", 2).
+               WillReturnResult(sqlmock.NewResult(0, 1))
+
+       result, err := ExecBatchInTxContext(ctx, tx, batchQuery, 
[][]any{{"user1", 1}, {"user2", 2}})
+       require.NoError(t, err)
+       require.Equal(t, BatchTransactionPending, 
result.Outcome.TransactionState)
+       require.Equal(t, BatchItemExecuted, result.Items[0].State)
+       require.Equal(t, BatchItemExecuted, result.Items[1].State)
+
+       // If the batch API committed the transaction internally,this statement 
would fail with sql.ErrTxDone
+       mock.ExpectExec(regexp.QuoteMeta(singleQuery)).
+               WithArgs(100, 10).
+               WillReturnResult(sqlmock.NewResult(0, 1))
+
+       _, err = tx.ExecContext(ctx, singleQuery, 100, 10)
+       require.NoError(t, err)
+
+       mock.ExpectCommit()
+       require.NoError(t, tx.Commit())
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchInTxContextDoesNotRollbackCallerTransaction(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       ctx := context.Background()
+       query := "UPDATE user SET name = ? WHERE id = ?"
+       execErr := errors.New("execute failed")
+
+       mock.ExpectBegin()
+       tx, err := db.BeginTx(ctx, nil)
+       require.NoError(t, err)
+
+       mock.ExpectExec(regexp.QuoteMeta(query)).
+               WithArgs("user1", 1).
+               WillReturnResult(sqlmock.NewResult(0, 1))
+
+       mock.ExpectExec(regexp.QuoteMeta(query)).
+               WithArgs("user2", 2).
+               WillReturnError(execErr)
+
+       result, err := ExecBatchInTxContext(ctx, tx, query, [][]any{{"user1", 
1}, {"user2", 2}})
+       require.ErrorIs(t, err, execErr)
+       require.Equal(t, BatchTransactionPending, 
result.Outcome.TransactionState)
+       require.Equal(t, 1, result.Outcome.FailedIndex)
+       require.Equal(t, BatchItemExecuted, result.Items[0].State)
+       require.Equal(t, BatchItemFailed, result.Items[1].State)
+
+       // The caller still owns the transaction
+       mock.ExpectRollback()
+
+       require.NoError(t, tx.Rollback())
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchContextEmptyBatchIsNoop(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       result, err := ExecBatchContext(context.Background(), db, "UPDATE user 
SET name = ? WHERE id = ?", nil)
+       require.NoError(t, err)
+       require.Empty(t, result.Items)
+       require.Equal(t, BatchTransactionNotStarted, 
result.Outcome.TransactionState)
+       // No Begin/Exec/Commit should happen.
+       require.NoError(t, mock.ExpectationsWereMet())
+}
+
+func TestExecBatchContextDoesNotLeakStateAfterFailure(t *testing.T) {
+       db, mock, err := sqlmock.New()
+       require.NoError(t, err)
+       defer db.Close()
+
+       ctx := context.Background()
+       query := "UPDATE user SET name = ? WHERE id = ?"
+
+       firstBatchErr := errors.New("first batch failed")
+
+       // Batch A.
+       mock.ExpectBegin()
+       mock.ExpectExec(regexp.QuoteMeta(query)).WithArgs("usera1", 
1).WillReturnResult(sqlmock.NewResult(0, 1))
+       mock.ExpectExec(regexp.QuoteMeta(query)).WithArgs("usera2", 
2).WillReturnError(firstBatchErr)
+       mock.ExpectRollback()
+
+       _, err = ExecBatchContext(ctx, db, query, [][]any{{"usera1", 1}, 
{"usera2", 2}})
+       require.ErrorIs(t, err, firstBatchErr)
+
+       // Batch B.
+       mock.ExpectBegin()
+       mock.ExpectExec(regexp.QuoteMeta(query)).WithArgs("userb1", 
3).WillReturnResult(sqlmock.NewResult(0, 1))
+       mock.ExpectExec(regexp.QuoteMeta(query)).WithArgs("userb2", 
4).WillReturnResult(sqlmock.NewResult(0, 1))
+
+       mock.ExpectCommit()
+
+       _, err = ExecBatchContext(ctx, db, query, [][]any{{"userb1", 3}, 
{"userb2", 4}})
+
+       require.NoError(t, err)
+       require.NoError(t, mock.ExpectationsWereMet())
+}
diff --git a/pkg/datasource/sql/conn.go b/pkg/datasource/sql/conn.go
index f3355c9e..45340155 100644
--- a/pkg/datasource/sql/conn.go
+++ b/pkg/datasource/sql/conn.go
@@ -37,6 +37,19 @@ type Conn struct {
        autoCommit bool
        dbName     string
        dbType     types.DBType
+       invalid    bool
+}
+
+func (c *Conn) invalidate() { c.invalid = true }
+
+func (c *Conn) IsValid() bool {
+       if c.invalid {
+               return false
+       }
+       if validator, ok := c.targetConn.(driver.Validator); ok {
+               return validator.IsValid()
+       }
+       return true
 }
 
 // ResetSession is called prior to executing a query on the connection
diff --git a/pkg/datasource/sql/conn_at.go b/pkg/datasource/sql/conn_at.go
index 196155a2..6fe27860 100644
--- a/pkg/datasource/sql/conn_at.go
+++ b/pkg/datasource/sql/conn_at.go
@@ -38,6 +38,17 @@ type rowsWithStmt struct {
        stmt driver.Stmt
 }
 
+// nonRetryableATError preserves the commit error without exposing the retry 
signal to database/sql.
+type nonRetryableATError struct{ cause error }
+
+func (e nonRetryableATError) Error() string { return e.cause.Error() }
+
+func (e nonRetryableATError) Is(target error) bool {
+       return target != driver.ErrBadConn && errors.Is(e.cause, target)
+}
+
+func (e nonRetryableATError) As(target any) bool { return errors.As(e.cause, 
target) }
+
 func (r *rowsWithStmt) Close() error {
        rowsErr := r.Rows.Close()
        stmtErr := r.stmt.Close()
@@ -275,6 +286,9 @@ func (c *ATConn) createTxAndExecIfNeeded(ctx 
context.Context, f func() (types.Ex
        // For ExecContext, commit the transaction if it was created
        if tx != nil {
                if err := tx.Commit(); err != nil {
+                       if errors.Is(err, driver.ErrBadConn) {
+                               return nil, nonRetryableATError{cause: err}
+                       }
                        return nil, err
                }
        }
diff --git a/pkg/datasource/sql/conn_at_test.go 
b/pkg/datasource/sql/conn_at_test.go
index 31f7f973..69a1c41d 100644
--- a/pkg/datasource/sql/conn_at_test.go
+++ b/pkg/datasource/sql/conn_at_test.go
@@ -39,7 +39,9 @@ import (
        atexec "seata.apache.org/seata-go/v2/pkg/datasource/sql/exec/at"
        "seata.apache.org/seata-go/v2/pkg/datasource/sql/mock"
        "seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+       "seata.apache.org/seata-go/v2/pkg/datasource/sql/undo"
        "seata.apache.org/seata-go/v2/pkg/protocol/branch"
+       "seata.apache.org/seata-go/v2/pkg/rm"
        "seata.apache.org/seata-go/v2/pkg/tm"
 )
 
@@ -975,3 +977,207 @@ func (mi *mockTxHook) BeforeRollback(tx *Tx) {
                mi.beforeRollback(tx)
        }
 }
+
+func TestATTxCommitFailuresBeforeLocalCommitRollback(t *testing.T) {
+       for _, failurePoint := range []string{"before commit", "branch 
registration", "undo manager lookup", "undo log flush"} {
+               t.Run(failurePoint, func(t *testing.T) {
+                       CleanTxHooks()
+                       t.Cleanup(CleanTxHooks)
+
+                       ctrl := gomock.NewController(t)
+                       defer ctrl.Finish()
+
+                       primaryErr := errors.New(failurePoint + " failed")
+                       localTx := mock.NewMockTestDriverTx(ctrl)
+                       rollbackCall := 
localTx.EXPECT().Rollback().Times(1).Return(nil)
+                       targetConn := mock.NewMockTestDriverConn(ctrl)
+                       txCtx := newATCommitFailureContext(types.DBTypeMySQL)
+                       expectedOutcome := atCommitOutcomeRolledBack
+                       var expectedRollbackErr, expectedReportErr error
+
+                       switch failurePoint {
+                       case "before commit":
+                               RegisterTxHook(&mockTxHook{beforeCommit: 
func(*Tx) error { return primaryErr }})
+                       case "branch registration":
+                               manager := newATCommitTestManager(t, ctrl)
+                               manager.EXPECT().BranchRegister(gomock.Any(), 
gomock.Any()).Times(1).Return(int64(0), primaryErr)
+                       case "undo manager lookup":
+                               expectedOutcome = atCommitOutcomeRollbackFailed
+                               expectedRollbackErr = errors.New("rollback 
cleanup failed")
+                               expectedReportErr = errors.New("branch failure 
report failed")
+                               rollbackCall.Return(expectedRollbackErr)
+                               txCtx.DBType = types.DBTypeOracle
+                               manager := newATCommitTestManager(t, ctrl)
+                               manager.EXPECT().BranchRegister(gomock.Any(), 
gomock.Any()).Times(1).Return(int64(123), nil)
+                               gomock.InOrder(rollbackCall, 
expectATCommitFailureReport(t, manager, expectedReportErr))
+                       case "undo log flush":
+                               previousUndoConfig := undo.UndoConfig
+                               undo.UndoConfig.LogSerialization = "json"
+                               t.Cleanup(func() { undo.UndoConfig = 
previousUndoConfig })
+
+                               manager := newATCommitTestManager(t, ctrl)
+                               manager.EXPECT().BranchRegister(gomock.Any(), 
gomock.Any()).Times(1).Return(int64(123), nil)
+                               
targetConn.EXPECT().PrepareContext(gomock.Any(), 
gomock.Any()).Times(1).Return(nil, primaryErr)
+                               gomock.InOrder(rollbackCall, 
expectATCommitFailureReport(t, manager, nil))
+                       }
+
+                       err := (&ATTx{tx: &Tx{
+                               conn:    &Conn{targetConn: targetConn},
+                               tranCtx: txCtx,
+                               target:  localTx,
+                       }}).Commit()
+
+                       if assert.Error(t, err) {
+                               var commitErr *atCommitError
+                               if assert.ErrorAs(t, err, &commitErr) {
+                                       assert.Equal(t, expectedOutcome, 
commitErr.outcome)
+                                       assert.Equal(t, expectedRollbackErr, 
commitErr.rollbackErr)
+                                       assert.Equal(t, expectedReportErr, 
commitErr.reportErr)
+                               }
+                               if failurePoint == "undo manager lookup" {
+                                       assert.Contains(t, err.Error(), "not 
found UndoLogManager")
+                               } else {
+                                       assert.ErrorIs(t, err, primaryErr)
+                               }
+                       }
+               })
+       }
+}
+
+func TestATTxCommitFailureDiscardsConnection(t *testing.T) {
+       for _, rollbackFailure := range []bool{true, false} {
+               name := "local commit failure"
+               if rollbackFailure {
+                       name = "rollback failure"
+               }
+               t.Run(name, func(t *testing.T) {
+                       CleanTxHooks()
+                       t.Cleanup(CleanTxHooks)
+                       exec.CleanCommonHook()
+                       t.Cleanup(exec.CleanCommonHook)
+
+                       ctrl := gomock.NewController(t)
+                       defer ctrl.Finish()
+
+                       var primaryErr error = errors.New(name)
+                       localTx := mock.NewMockTestDriverTx(ctrl)
+                       var cleanupErr error
+                       expectedOutcome := atCommitOutcomeCommitUnknown
+                       var commitCall *gomock.Call
+                       if rollbackFailure {
+                               cleanupErr = errors.New("rollback cleanup 
failed")
+                               expectedOutcome = atCommitOutcomeRollbackFailed
+                               RegisterTxHook(&mockTxHook{beforeCommit: 
func(*Tx) error { return primaryErr }})
+                               
localTx.EXPECT().Rollback().Times(1).Return(cleanupErr)
+                       } else {
+                               primaryErr = driver.ErrBadConn
+                               RegisterTxHook(&mockTxHook{beforeCommit: 
func(tx *Tx) error {
+                                       tx.tranCtx.BranchID = 123
+                                       return nil
+                               }})
+                               commitCall = 
localTx.EXPECT().Commit().Times(1).Return(primaryErr)
+                               localTx.EXPECT().Rollback().Times(0)
+                       }
+
+                       targetConn := mock.NewMockTestDriverConn(ctrl)
+                       targetConn.EXPECT().BeginTx(gomock.Any(), 
gomock.Any()).Times(1).Return(localTx, nil)
+                       query := "SELECT 1"
+                       var args []any
+                       if rollbackFailure {
+                               previousTableCache := 
datasource.GetTableCache(types.DBTypeMySQL)
+                               t.Cleanup(func() {
+                                       
datasource.RegisterTableCache(types.DBTypeMySQL, previousTableCache)
+                               })
+                               
datasource.RegisterTableCache(types.DBTypeMySQL, batchATTableCache{})
+
+                               query = "UPDATE account SET balance = ? WHERE 
id = ?"
+                               args = []any{int64(110), int64(1)}
+                               expectBatchATImageQueries(t, targetConn, 
[][]driver.Value{
+                                       {int64(1), int64(100)}, {int64(1), 
int64(110)},
+                               })
+                       }
+                       targetConn.EXPECT().ExecContext(gomock.Any(), query, 
gomock.Any()).Times(1).Return(driver.ResultNoRows, nil)
+                       var closeCount atomic.Int32
+                       
targetConn.EXPECT().Close().AnyTimes().DoAndReturn(func() error {
+                               closeCount.Add(1)
+                               return nil
+                       })
+
+                       txCtx := types.NewTxCtx()
+                       txCtx.XID = "test-xid"
+                       txCtx.TransactionMode = types.ATMode
+                       txCtx.DBType = types.DBTypeMySQL
+                       if !rollbackFailure {
+                               manager := newATCommitTestManager(t, ctrl)
+                               gomock.InOrder(commitCall, 
expectATCommitFailureReport(t, manager, nil))
+                       }
+                       conn := &Conn{
+                               res:        &DBResource{dbType: 
types.DBTypeMySQL, resourceID: "test-resource"},
+                               txCtx:      txCtx,
+                               targetConn: targetConn,
+                               autoCommit: true,
+                       }
+                       connector := mock.NewMockTestDriverConnector(ctrl)
+                       
connector.EXPECT().Connect(gomock.Any()).Times(1).Return(&ATConn{Conn: conn}, 
nil)
+                       db := sql.OpenDB(connector)
+                       defer db.Close()
+
+                       ctx := tm.InitSeataContext(context.Background())
+                       tm.SetXID(ctx, "test-xid")
+                       _, err := db.ExecContext(ctx, query, args...)
+
+                       assert.Contains(t, err.Error(), primaryErr.Error())
+                       assert.NotErrorIs(t, err, driver.ErrBadConn)
+                       var commitErr *atCommitError
+                       if assert.ErrorAs(t, err, &commitErr) {
+                               assert.Equal(t, primaryErr, commitErr.cause)
+                               assert.Equal(t, expectedOutcome, 
commitErr.outcome)
+                               assert.Equal(t, cleanupErr, 
commitErr.rollbackErr)
+                       }
+                       if cleanupErr != nil {
+                               assert.Contains(t, err.Error(), 
cleanupErr.Error())
+                       }
+                       validator, ok := any(conn).(driver.Validator)
+                       assert.True(t, ok)
+                       if ok {
+                               assert.False(t, validator.IsValid())
+                       }
+                       assert.Equal(t, int32(1), closeCount.Load())
+               })
+       }
+}
+
+func newATCommitFailureContext(dbType types.DBType) *types.TransactionContext {
+       txCtx := types.NewTxCtx()
+       txCtx.XID = "test-xid"
+       txCtx.ResourceID = "test-resource"
+       txCtx.TransactionMode = types.ATMode
+       txCtx.DBType = dbType
+       txCtx.LockKeys["test_table:1"] = struct{}{}
+       txCtx.RoundImages.AppendBeofreImage(&types.RecordImage{
+               TableName: "test_table",
+               SQLType:   types.SQLTypeUpdate,
+               Rows:      []types.RowImage{{}},
+       })
+       return txCtx
+}
+
+func newATCommitTestManager(t *testing.T, ctrl *gomock.Controller) 
*mock.MockDataSourceManager {
+       manager := mock.NewMockDataSourceManager(ctrl)
+       manager.SetBranchType(branch.BranchTypeAT)
+       registerResourceManagerForTest(t, manager)
+       return manager
+}
+
+func expectATCommitFailureReport(t *testing.T, manager 
*mock.MockDataSourceManager, reportErr error) *gomock.Call {
+       times := 1
+       if reportErr != nil {
+               times = 5
+       }
+       return manager.EXPECT().BranchReport(gomock.Any(), 
gomock.Any()).Times(times).DoAndReturn(
+               func(_ context.Context, param rm.BranchReportParam) error {
+                       assert.EqualValues(t, 
branch.BranchStatusPhaseoneFailed, param.Status)
+                       return reportErr
+               },
+       )
+}
diff --git a/pkg/datasource/sql/tx_at.go b/pkg/datasource/sql/tx_at.go
index 7ee107b0..36174a32 100644
--- a/pkg/datasource/sql/tx_at.go
+++ b/pkg/datasource/sql/tx_at.go
@@ -18,11 +18,48 @@
 package sql
 
 import (
-       "github.com/pkg/errors"
+       "fmt"
 
        "seata.apache.org/seata-go/v2/pkg/datasource/sql/undo"
 )
 
+type localCommitStage uint8
+
+const (
+       localCommitNotStarted localCommitStage = iota
+       localCommitInvoked
+       localCommitSucceeded
+)
+
+type atCommitOutcome uint8
+
+const (
+       atCommitOutcomeRolledBack atCommitOutcome = iota
+       atCommitOutcomeRollbackFailed
+       atCommitOutcomeCommitUnknown
+       atCommitOutcomeCommitted
+)
+
+type atCommitError struct {
+       cause       error
+       rollbackErr error
+       reportErr   error
+       outcome     atCommitOutcome
+}
+
+func (e *atCommitError) Error() string {
+       message := fmt.Sprintf("AT commit failed: %v", e.cause)
+       if e.rollbackErr != nil {
+               message += fmt.Sprintf("; rollback failed: %v", e.rollbackErr)
+       }
+       if e.reportErr != nil {
+               message += fmt.Sprintf("; branch failure report failed: %v", 
e.reportErr)
+       }
+       return message
+}
+
+func (e *atCommitError) Unwrap() error { return e.cause }
+
 // ATTx
 type ATTx struct {
        tx *Tx
@@ -33,10 +70,11 @@ type ATTx struct {
 // case 2. not need flush undolog, is XA mode, do local transaction commit
 // case 3. need run AT transaction
 func (tx *ATTx) Commit() error {
-       if err := tx.tx.beforeCommit(); err != nil {
-               return err
+       stage, err := tx.doCommit()
+       if err == nil {
+               return nil
        }
-       return tx.commitOnAT()
+       return tx.finishCommitFailure(stage, err)
 }
 
 func (tx *ATTx) Rollback() error {
@@ -53,32 +91,57 @@ func (tx *ATTx) Rollback() error {
        return err
 }
 
-// commitOnAT
-func (tx *ATTx) commitOnAT() error {
+func (tx *ATTx) doCommit() (localCommitStage, error) {
        originTx := tx.tx
+       stage := localCommitNotStarted
+
+       if err := originTx.beforeCommit(); err != nil {
+               return stage, err
+       }
+
        if err := originTx.register(originTx.tranCtx); err != nil {
-               return err
+               return stage, err
        }
 
        undoLogMgr, err := undo.GetUndoLogManager(originTx.tranCtx.DBType)
        if err != nil {
-               return err
+               return stage, err
        }
 
        if err = undoLogMgr.FlushUndoLog(originTx.tranCtx, 
originTx.conn.targetConn); err != nil {
-               if rerr := originTx.report(false); rerr != nil {
-                       return errors.WithStack(rerr)
-               }
-               return errors.WithStack(err)
+               return stage, err
        }
 
+       stage = localCommitInvoked
        if err := originTx.commitOnLocal(); err != nil {
-               if rerr := originTx.report(false); rerr != nil {
-                       return errors.WithStack(rerr)
-               }
-               return errors.WithStack(err)
+               return stage, err
        }
 
+       stage = localCommitSucceeded
        originTx.report(true)
-       return nil
+       return stage, nil
+}
+
+func (tx *ATTx) finishCommitFailure(stage localCommitStage, cause error) error 
{
+       originTx := tx.tx
+       commitErr := &atCommitError{cause: cause, outcome: 
atCommitOutcomeRolledBack}
+
+       switch stage {
+       case localCommitNotStarted:
+               if err := originTx.Rollback(); err != nil {
+                       commitErr.rollbackErr = err
+                       commitErr.outcome = atCommitOutcomeRollbackFailed
+                       originTx.conn.invalidate()
+               }
+       case localCommitInvoked:
+               commitErr.outcome = atCommitOutcomeCommitUnknown
+               originTx.conn.invalidate()
+       case localCommitSucceeded:
+               commitErr.outcome = atCommitOutcomeCommitted
+       }
+       if stage != localCommitSucceeded && 
originTx.tranCtx.IsBranchRegistered() {
+               commitErr.reportErr = originTx.report(false)
+       }
+
+       return commitErr
 }


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to