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 6a483372 feat:support MySQL multi-statement DML in AT mode (#1140)
6a483372 is described below
commit 6a483372357dbf53d9aefdaf9191f9ba50c6560f
Author: Mochimia <[email protected]>
AuthorDate: Sat Aug 22 23:30:36 2026 +0800
feat:support MySQL multi-statement DML in AT mode (#1140)
---
pkg/datasource/sql/conn_at.go | 74 +-
pkg/datasource/sql/conn_at_test.go | 97 +++
pkg/datasource/sql/exec/at/at_executor.go | 1 +
pkg/datasource/sql/exec/at/at_executor_test.go | 7 +-
pkg/datasource/sql/exec/at/base_executor.go | 1 +
.../sql/exec/at/multi_delete_executor.go | 26 +-
pkg/datasource/sql/exec/at/multi_execution_plan.go | 194 ++++++
.../sql/exec/at/multi_execution_plan_test.go | 278 ++++++++
pkg/datasource/sql/exec/at/multi_executor.go | 101 ++-
pkg/datasource/sql/exec/at/multi_executor_test.go | 67 ++
.../sql/exec/at/multi_sequential_executor.go | 262 +++++++
.../sql/exec/at/multi_sequential_executor_test.go | 771 +++++++++++++++++++++
pkg/datasource/sql/exec/at/multi_update_excutor.go | 66 +-
pkg/datasource/sql/exec/at/update_join_executor.go | 2 +-
pkg/datasource/sql/exec/hook.go | 10 +
pkg/datasource/sql/util/ctxutil.go | 89 +++
pkg/datasource/sql/util/ctxutil_test.go | 164 +++++
17 files changed, 2123 insertions(+), 87 deletions(-)
diff --git a/pkg/datasource/sql/conn_at.go b/pkg/datasource/sql/conn_at.go
index 4d3bc66c..196155a2 100644
--- a/pkg/datasource/sql/conn_at.go
+++ b/pkg/datasource/sql/conn_at.go
@@ -25,7 +25,9 @@ import (
"strings"
"seata.apache.org/seata-go/v2/pkg/datasource/sql/exec"
+ sqlparser "seata.apache.org/seata-go/v2/pkg/datasource/sql/parser"
"seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/util"
"seata.apache.org/seata-go/v2/pkg/tm"
"seata.apache.org/seata-go/v2/pkg/util/log"
)
@@ -51,7 +53,45 @@ type ATConn struct {
*Conn
}
+var (
+ errATPreparedMultiSQLUnsupported = errors.New("seata AT: prepared
multi-SQL is unsupported; use Exec or ExecContext")
+ parseATPreparedSQL = sqlparser.DoParser
+)
+
+func rejectATPreparedMultiSQL(dbType types.DBType, query string) error {
+ if dbType != types.DBTypeMySQL || !strings.Contains(query, ";") {
+ return nil
+ }
+
+ parseCtx, err := parseATPreparedSQL(query)
+ if err != nil || parseCtx == nil {
+ // This is a best-effort Prepare guard, not the AT
execution-time
+ // validation boundary. Preserve Prepare compatibility for SQL
that
+ // the parser cannot handle; BuildExecutor parses it again
before
+ // AT execution.
+ return nil
+ }
+
+ if len(parseCtx.MultiStmt) > 1 {
+ return errATPreparedMultiSQLUnsupported
+ }
+
+ return nil
+}
+
+func (c *ATConn) Prepare(query string) (driver.Stmt, error) {
+ if err := rejectATPreparedMultiSQL(c.dbType, query); err != nil {
+ return nil, err
+ }
+
+ return c.Conn.Prepare(query)
+}
+
func (c *ATConn) PrepareContext(ctx context.Context, query string)
(driver.Stmt, error) {
+ if err := rejectATPreparedMultiSQL(c.dbType, query); err != nil {
+ return nil, err
+ }
+
if c.createOnceTxContext(ctx) {
defer func() {
c.txCtx = types.NewTxCtx()
@@ -78,38 +118,12 @@ func (c *ATConn) ExecContext(ctx context.Context, query
string, args []driver.Na
ret, err := executor.ExecWithNamedValue(ctx, execCtx,
func(ctx context.Context, query string, args
[]driver.NamedValue) (types.ExecResult, error) {
- ret, err := c.Conn.ExecContext(ctx, query, args)
- if err == nil {
- return
types.NewResult(types.WithResult(ret)), nil
- }
-
- // If skip fast-path error, fallback to
prepared statement
- if strings.Contains(err.Error(), "skip
fast-path") {
- stmt, prepErr := c.Conn.Prepare(query)
- if prepErr != nil {
- return nil, prepErr
- }
- defer stmt.Close()
-
- var result driver.Result
- if stmtExecCtx, ok :=
stmt.(driver.StmtExecContext); ok {
- result, err =
stmtExecCtx.ExecContext(ctx, args)
- } else {
- dargs := make([]driver.Value,
len(args))
- for i, arg := range args {
- dargs[i] = arg.Value
- }
- result, err = stmt.Exec(dargs)
- }
-
- if err != nil {
- return nil, err
- }
-
- return
types.NewResult(types.WithResult(result)), nil
+ result, err :=
util.CtxDriverExecWithPrepareFallback(ctx, c.targetConn, query, args)
+ if err != nil {
+ return nil, err
}
- return nil, err
+ return
types.NewResult(types.WithResult(result)), nil
})
if err != nil {
diff --git a/pkg/datasource/sql/conn_at_test.go
b/pkg/datasource/sql/conn_at_test.go
index e5569d1c..31f7f973 100644
--- a/pkg/datasource/sql/conn_at_test.go
+++ b/pkg/datasource/sql/conn_at_test.go
@@ -48,6 +48,103 @@ func TestMain(m *testing.M) {
m.Run()
}
+func TestATConnRejectsPreparedMultiSQL(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+
+ conn := &ATConn{Conn: &Conn{targetConn: targetConn, dbType:
types.DBTypeMySQL}}
+
+ queries := []string{"INSERT INTO t_user(id) VALUES (?);" + "INSERT INTO
t_user(id) VALUES (?)",
+ "UPDATE t_user SET name = ? WHERE id = ?;" + "DELETE FROM
t_user_log WHERE user_id = ?",
+ }
+
+ for _, query := range queries {
+ stmt, err := conn.Prepare(query)
+ assert.Nil(t, stmt)
+ assert.ErrorIs(t, err, errATPreparedMultiSQLUnsupported)
+
+ stmt, err = conn.PrepareContext(context.Background(), query)
+ assert.Nil(t, stmt)
+ assert.ErrorIs(t, err, errATPreparedMultiSQLUnsupported)
+ }
+}
+
+func TestRejectATPreparedMultiSQLAllowsSingleStatement(t *testing.T) {
+ err := rejectATPreparedMultiSQL(types.DBTypeMySQL, "UPDATE t_user SET
name = ? WHERE id = ?")
+
+ assert.NoError(t, err)
+}
+
+func TestRejectATPreparedMultiSQLSkipsParserWithoutSemicolon(t *testing.T) {
+ originalParseATPreparedSQL := parseATPreparedSQL
+ t.Cleanup(func() { parseATPreparedSQL = originalParseATPreparedSQL })
+
+ parseCalls := 0
+ parseATPreparedSQL = func(string) (*types.ParseContext, error) {
+ parseCalls++
+ return nil, errors.New("unexpected parser call")
+ }
+
+ err := rejectATPreparedMultiSQL(types.DBTypeMySQL, "UPDATE t_user SET
name = ? WHERE id = ?")
+
+ assert.NoError(t, err)
+ assert.Zero(t, parseCalls)
+}
+
+func TestRejectATPreparedMultiSQLAllowsParserFailure(t *testing.T) {
+ originalParseATPreparedSQL := parseATPreparedSQL
+ t.Cleanup(func() { parseATPreparedSQL = originalParseATPreparedSQL })
+
+ parserError := errors.New("unsupported SQL syntax")
+ tests := []struct {
+ name string
+ parseCtx *types.ParseContext
+ err error
+ }{
+ {
+ name: "parser returns error",
+ err: parserError,
+ },
+ {
+ name: "parser returns nil context",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ parseCalls := 0
+ parseATPreparedSQL = func(string) (*types.ParseContext,
error) {
+ parseCalls++
+ return tt.parseCtx, tt.err
+ }
+
+ err := rejectATPreparedMultiSQL(types.DBTypeMySQL,
"UPDATE t_user SET name = ?; unsupported syntax")
+
+ assert.NoError(t, err)
+ assert.Equal(t, 1, parseCalls)
+ })
+ }
+}
+
+func TestATConnAllowsPreparedMultiSQLForPostgreSQL(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+ targetStmt := mock.NewMockTestDriverStmt(ctrl)
+ query := "UPDATE t_user SET name = $1 WHERE id = $2;DELETE FROM
t_user_log WHERE user_id = $3"
+ conn := &ATConn{Conn: &Conn{targetConn: targetConn, dbType:
types.DBTypePostgreSQL}}
+
+ targetConn.EXPECT().Prepare(query).Return(targetStmt, nil)
+ targetConn.EXPECT().PrepareContext(gomock.Any(),
query).Return(targetStmt, nil)
+
+ stmt, err := conn.Prepare(query)
+ assert.NotNil(t, stmt)
+ assert.NoError(t, err)
+
+ stmt, err = conn.PrepareContext(context.Background(), query)
+ assert.NotNil(t, stmt)
+ assert.NoError(t, err)
+}
+
type postgresMockRows struct {
columns []string
data [][]driver.Value
diff --git a/pkg/datasource/sql/exec/at/at_executor.go
b/pkg/datasource/sql/exec/at/at_executor.go
index ca6b1ea3..13f4c7ad 100644
--- a/pkg/datasource/sql/exec/at/at_executor.go
+++ b/pkg/datasource/sql/exec/at/at_executor.go
@@ -31,6 +31,7 @@ import (
var (
parseSQLQuery = parser.DoParser
isGlobalTx = tm.IsGlobalTx
+ hooksForSQLType = exec.HooksForSQLType
newPlainExecutor = NewPlainExecutor
newInsertExecutor = NewInsertExecutor
newUpdateExecutor = NewUpdateExecutor
diff --git a/pkg/datasource/sql/exec/at/at_executor_test.go
b/pkg/datasource/sql/exec/at/at_executor_test.go
index 54b9a584..f3b8b758 100644
--- a/pkg/datasource/sql/exec/at/at_executor_test.go
+++ b/pkg/datasource/sql/exec/at/at_executor_test.go
@@ -407,6 +407,11 @@ func TestATExecutor_ExecWithValue_ParserError(t
*testing.T) {
func TestATExecutors_ExecContext_BeforeHookError(t *testing.T) {
beforeErr := fmt.Errorf("before hook error")
+ multiParseCtx, err := parseSQLQuery("INSERT INTO t_user(id) VALUES
(1);" + "INSERT INTO t_user(id) VALUES (2)")
+ if !assert.NoError(t, err) {
+ return
+ }
+
tests := []struct {
name string
newExecutor func(hooks []exec.SQLHook) executor
@@ -438,7 +443,7 @@ func TestATExecutors_ExecContext_BeforeHookError(t
*testing.T) {
{
name: "multi",
newExecutor: func(hooks []exec.SQLHook) executor {
- return &multiExecutor{baseExecutor:
baseExecutor{hooks: hooks}, execContext: &types.ExecContext{}}
+ return &multiExecutor{baseExecutor:
baseExecutor{hooks: hooks}, parserCtx: multiParseCtx, execContext:
&types.ExecContext{DBType: types.DBTypeMySQL}}
},
},
{
diff --git a/pkg/datasource/sql/exec/at/base_executor.go
b/pkg/datasource/sql/exec/at/base_executor.go
index 9479df43..28918c9d 100644
--- a/pkg/datasource/sql/exec/at/base_executor.go
+++ b/pkg/datasource/sql/exec/at/base_executor.go
@@ -586,6 +586,7 @@ func (b *baseExecutor) buildLockKey(records
*types.RecordImage, meta types.Table
func (b *baseExecutor) rowsPrepare(ctx context.Context, conn driver.Conn,
selectSQL string, selectArgs []driver.NamedValue) (driver.Rows, error) {
var queryer driver.Queryer
+ var rows driver.Rows
queryerContext, ok := conn.(driver.QueryerContext)
if !ok {
diff --git a/pkg/datasource/sql/exec/at/multi_delete_executor.go
b/pkg/datasource/sql/exec/at/multi_delete_executor.go
index 3f2ad4ec..2c55c360 100644
--- a/pkg/datasource/sql/exec/at/multi_delete_executor.go
+++ b/pkg/datasource/sql/exec/at/multi_delete_executor.go
@@ -88,26 +88,20 @@ func (m *multiDeleteExecutor) beforeImage(ctx
context.Context) ([]*types.RecordI
records []*types.RecordImage
)
- queryerCtx, ok := m.execContext.Conn.(driver.QueryerContext)
- var queryer driver.Queryer
- if !ok {
- queryer, ok = m.execContext.Conn.(driver.Queryer)
- }
- if !ok {
- log.Errorf("target conn should been driver.QueryerContext or
driver.Queryer")
- return nil, fmt.Errorf("invalid conn")
+ rowsi, err = util.CtxDriverQueryWithPrepareFallback(ctx,
m.execContext.Conn, selectSQL, args)
+ if err != nil {
+ log.Errorf("aggregate delete image query failed: %+v", err)
+ return nil, err
}
-
- rowsi, err = util.CtxDriverQuery(ctx, queryerCtx, queryer, selectSQL,
args)
defer func() {
- if rowsi != nil {
- rowsi.Close()
+ if rowsi == nil {
+ return
+ }
+
+ if closeErr := rowsi.Close(); closeErr != nil {
+ log.Errorf("rows close fail,err: %v", closeErr)
}
}()
- if err != nil {
- log.Errorf("ctx driver query: %+v", err)
- return nil, err
- }
tableName, err := m.getFromTableInSQL()
if err != nil {
diff --git a/pkg/datasource/sql/exec/at/multi_execution_plan.go
b/pkg/datasource/sql/exec/at/multi_execution_plan.go
new file mode 100644
index 00000000..15468c1e
--- /dev/null
+++ b/pkg/datasource/sql/exec/at/multi_execution_plan.go
@@ -0,0 +1,194 @@
+/*
+ * 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 at
+
+import (
+ "errors"
+ "fmt"
+
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+)
+
+var (
+ errInvalidMultiSQL = errors.New("invalid multi SQL")
+ errUnsupportedMultiSQL = errors.New("unsupported multi SQL")
+)
+
+// multiExecutionPlan contains the validated statements in their original order
+// and records whether they can use the existing aggregate execution path.
+type multiExecutionPlan struct {
+ statements []*types.ParseContext
+ useAggregatePath bool
+}
+
+// buildMultiExecutionPlan validates all statements before any before-image
query
+// or business SQL execution occurs.
+func buildMultiExecutionPlan(parseCtx *types.ParseContext, dbType
types.DBType) (*multiExecutionPlan, error) {
+ if parseCtx == nil {
+ return nil, fmt.Errorf("%w: parse context", errInvalidMultiSQL)
+ }
+
+ if len(parseCtx.MultiStmt) < 2 {
+ return nil, fmt.Errorf(
+ "%w: expected at least two statements, got %d",
+ errInvalidMultiSQL, len(parseCtx.MultiStmt))
+ }
+
+ // Multi-SQL AT execution is currently enabled only for MySQL.
+ if effectiveDBType(dbType) != types.DBTypeMySQL {
+ return nil, fmt.Errorf(
+ "%w: database type %v is not supported",
+ errUnsupportedMultiSQL, dbType,
+ )
+ }
+
+ statements := append([]*types.ParseContext(nil), parseCtx.MultiStmt...)
+ tableNames := make([]string, len(statements))
+
+ for index, statementCtx := range statements {
+ tableName, err := validateMultiStatement(index, statementCtx)
+ if err != nil {
+ return nil, err
+ }
+ tableNames[index] = tableName
+ }
+
+ return &multiExecutionPlan{
+ statements: statements,
+ useAggregatePath: canUseAggregateFastPath(statements,
tableNames),
+ }, nil
+}
+
+// validateMultiStatement validates one parsed DML statement.
+//
+// The table name is returned temporarily for aggregate-path detection.
+func validateMultiStatement(index int, parseCtx *types.ParseContext) (string,
error) {
+ if parseCtx == nil {
+ return "", fmt.Errorf("%w: statement %d parse context is nil",
errInvalidMultiSQL, index)
+ }
+
+ switch parseCtx.ExecutorType {
+ case types.InsertExecutor:
+ if parseCtx.InsertStmt == nil {
+ return "", fmt.Errorf("%w: statement %d is marked as
INSERT but has no INSERT AST", errInvalidMultiSQL, index)
+ }
+
+ case types.UpdateExecutor:
+ if parseCtx.UpdateStmt == nil {
+ return "", fmt.Errorf("%w: statement %d is marked as
UPDATE but has no UPDATE AST", errInvalidMultiSQL, index)
+ }
+
+ updateStmt := parseCtx.UpdateStmt
+ if updateStmt.TableRefs == nil ||
updateStmt.TableRefs.TableRefs == nil {
+ return "", fmt.Errorf("%w: statement %d has invalid
UPDATE table references", errInvalidMultiSQL, index)
+ }
+
+ if updateStmt.TableRefs.TableRefs.Right != nil {
+ return "", fmt.Errorf("%w: statement %d uses UPDATE
JOIN", errUnsupportedMultiSQL, index)
+ }
+
+ case types.DeleteExecutor:
+ if parseCtx.DeleteStmt == nil {
+ return "", fmt.Errorf("%w: statement %d is marked as
DELETE but has no DELETE AST", errInvalidMultiSQL, index)
+ }
+
+ if parseCtx.DeleteStmt.IsMultiTable {
+ return "", fmt.Errorf("%w: statement %d uses
multi-table DELETE", errUnsupportedMultiSQL, index)
+ }
+
+ default:
+ return "", fmt.Errorf("%w: statement %d uses executor type %v",
errUnsupportedMultiSQL, index, parseCtx.ExecutorType)
+ }
+
+ tableName, err := parseCtx.GetTableName()
+ if err != nil {
+ return "", fmt.Errorf("%w: get table name for statement %d:
%w", errInvalidMultiSQL, index, err)
+ }
+
+ return tableName, nil
+}
+
+// canUseAggregateFastPath reports whether all statements satisfy the static
+// contract of the existing aggregate UPDATE/DELETE executors.
+//
+// Aggregate execution is an optimization. Any statement that cannot be proven
+// safe must use the sequential single-statement executors.
+func canUseAggregateFastPath(statements []*types.ParseContext, tableNames
[]string) bool {
+ if len(statements) < 2 || len(statements) != len(tableNames) {
+ return false
+ }
+
+ firstStatement := statements[0]
+ if firstStatement == nil {
+ return false
+ }
+
+ if firstStatement.ExecutorType != types.UpdateExecutor &&
firstStatement.ExecutorType != types.DeleteExecutor {
+ return false
+ }
+
+ firstTableName := tableNames[0]
+
+ for index, statementCtx := range statements {
+ if statementCtx == nil {
+ return false
+ }
+
+ if statementCtx.ExecutorType != firstStatement.ExecutorType {
+ return false
+ }
+
+ if tableNames[index] != firstTableName {
+ return false
+ }
+
+ statementNode, err := getStatementNode(statementCtx)
+ if err != nil {
+ return false
+ }
+
+ parameterCount, err := countStatementParameters(statementNode)
+ if err != nil || parameterCount != 0 {
+ return false
+ }
+
+ switch statementCtx.ExecutorType {
+ case types.UpdateExecutor:
+ updateStmt := statementCtx.UpdateStmt
+ if updateStmt == nil || updateStmt.Where == nil ||
updateStmt.Limit != nil || updateStmt.Order != nil {
+ return false
+ }
+
+ // UPDATE JOIN uses the single-statement update-join
executor.
+ if updateStmt.TableRefs == nil ||
updateStmt.TableRefs.TableRefs == nil || updateStmt.TableRefs.TableRefs.Right
!= nil {
+ return false
+ }
+
+ case types.DeleteExecutor:
+ deleteStmt := statementCtx.DeleteStmt
+ if deleteStmt == nil || deleteStmt.IsMultiTable ||
deleteStmt.Where == nil || deleteStmt.Limit != nil || deleteStmt.Order != nil {
+ return false
+ }
+
+ default:
+ return false
+ }
+ }
+
+ return true
+}
diff --git a/pkg/datasource/sql/exec/at/multi_execution_plan_test.go
b/pkg/datasource/sql/exec/at/multi_execution_plan_test.go
new file mode 100644
index 00000000..47f9421e
--- /dev/null
+++ b/pkg/datasource/sql/exec/at/multi_execution_plan_test.go
@@ -0,0 +1,278 @@
+/*
+ * 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 at
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/assert"
+
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/parser"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+)
+
+func TestBuildMultiExecutionPlan(t *testing.T) {
+ tests := []struct {
+ name string
+ sourceQuery string
+ expectStatementSize int
+ expectAggregatePath bool
+ }{
+ {
+ name: "literal same table updates use aggregate path",
+ sourceQuery: "UPDATE t_user SET name = 'user1' WHERE id
= 1;" +
+ "UPDATE t_user SET age = 18 WHERE id = 2",
+ expectStatementSize: 2,
+ expectAggregatePath: true,
+ },
+ {
+ name: "parameterized same table updates use sequential
path",
+ sourceQuery: "UPDATE t_user SET name = ? WHERE id = ?;"
+
+ "UPDATE t_user SET age = ? WHERE id = ?",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "updates without where use sequential path",
+ sourceQuery: "UPDATE t_user SET name = 'user1';" +
+ "UPDATE t_user SET age = 18",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "literal same table deletes use aggregate path",
+ sourceQuery: "DELETE FROM t_user WHERE id = 1;" +
+ "DELETE FROM t_user WHERE id = 2",
+ expectStatementSize: 2,
+ expectAggregatePath: true,
+ },
+ {
+ name: "parameterized same table deletes use sequential
path",
+ sourceQuery: "DELETE FROM t_user WHERE id = ?;" +
+ "DELETE FROM t_user WHERE id = ?",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "deletes without where use sequential path",
+ sourceQuery: "DELETE FROM t_user;" +
+ "DELETE FROM t_user",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "multiple inserts use sequential path",
+ sourceQuery: "INSERT INTO t_user(id, name) VALUES (?,
?);" +
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "literal update with limit uses sequential path",
+ sourceQuery: "UPDATE t_user SET status = 1 " +
+ "WHERE status = 0 LIMIT 10;" +
+ "UPDATE t_user SET status = 2 " +
+ "WHERE id = 1",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "updates on different tables use sequential path",
+ sourceQuery: "UPDATE t_user SET name = ? WHERE id = ?;"
+
+ "UPDATE t_account SET balance = ? WHERE id = ?",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "deletes on different tables use sequential path",
+ sourceQuery: "DELETE FROM t_user WHERE id = ?;" +
+ "DELETE FROM t_user_log WHERE user_id = ?",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "update with limit uses sequential path",
+ sourceQuery: "UPDATE t_user SET status = ? WHERE status
= ? LIMIT ?;" +
+ "UPDATE t_user SET status = ? WHERE id = ?",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "update with order by uses sequential path",
+ sourceQuery: "UPDATE t_user SET status = ? WHERE status
= ? ORDER BY id;" +
+ "UPDATE t_user SET status = ? WHERE id = ?",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "delete with limit uses sequential path",
+ sourceQuery: "DELETE FROM t_user WHERE status = ? LIMIT
?;" +
+ "DELETE FROM t_user WHERE id = ?",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ {
+ name: "delete with order by uses sequential path",
+ sourceQuery: "DELETE FROM t_user WHERE status = ? ORDER
BY id;" +
+ "DELETE FROM t_user WHERE id = ?",
+ expectStatementSize: 2,
+ expectAggregatePath: false,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ parseCtx, err := parser.DoParser(tt.sourceQuery)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ plan, err := buildMultiExecutionPlan(parseCtx,
types.DBTypeMySQL)
+ if !assert.NoError(t, err) || !assert.NotNil(t, plan) {
+ return
+ }
+
+ assert.Len(t, plan.statements, tt.expectStatementSize)
+ assert.Equal(t, tt.expectAggregatePath,
plan.useAggregatePath)
+
+ for index := range parseCtx.MultiStmt {
+ assert.Same(t, parseCtx.MultiStmt[index],
plan.statements[index])
+ }
+ })
+ }
+}
+
+func TestBuildMultiExecutionPlanError(t *testing.T) {
+ tests := []struct {
+ name string
+ parseCtx *types.ParseContext
+ dbType types.DBType
+ expectError error
+ }{
+ {
+ name: "nil parse context",
+ parseCtx: nil,
+ dbType: types.DBTypeMySQL,
+ expectError: errInvalidMultiSQL,
+ },
+ {
+ name: "single statement",
+ parseCtx: mustParseMultiExecutionPlanTestSQL(t,
"UPDATE t_user SET name = ? WHERE id = ?"),
+ dbType: types.DBTypeMySQL,
+ expectError: errInvalidMultiSQL,
+ },
+ {
+ name: "unsupported database type",
+ parseCtx: mustParseMultiExecutionPlanTestSQL(
+ t,
+ "UPDATE t_user SET name = ? WHERE id = ?;"+
+ "UPDATE t_user SET age = ? WHERE id =
?",
+ ),
+ dbType: types.DBTypePostgreSQL,
+ expectError: errUnsupportedMultiSQL,
+ },
+ {
+ name: "unsupported select statement",
+ parseCtx: mustParseMultiExecutionPlanTestSQL(
+ t,
+ "SELECT * FROM t_user WHERE id = ?;"+
+ "UPDATE t_user SET name = ? WHERE id =
?",
+ ),
+ dbType: types.DBTypeMySQL,
+ expectError: errUnsupportedMultiSQL,
+ },
+ {
+ name: "nil child parse context",
+ parseCtx: &types.ParseContext{
+ SQLType: types.SQLTypeMulti,
+ ExecutorType: types.MultiExecutor,
+ MultiStmt: []*types.ParseContext{
+ mustParseMultiExecutionPlanTestSQL(t,
"UPDATE t_user SET name = ? WHERE id = ?"),
+ nil,
+ },
+ },
+ dbType: types.DBTypeMySQL,
+ expectError: errInvalidMultiSQL,
+ },
+ {
+ name: "update executor without update AST",
+ parseCtx: &types.ParseContext{
+ SQLType: types.SQLTypeMulti,
+ ExecutorType: types.MultiExecutor,
+ MultiStmt: []*types.ParseContext{
+ {
+ SQLType:
types.SQLTypeUpdate,
+ ExecutorType:
types.UpdateExecutor,
+ },
+ mustParseMultiExecutionPlanTestSQL(t,
"UPDATE t_user SET age = ? WHERE id = ?"),
+ },
+ },
+ dbType: types.DBTypeMySQL,
+ expectError: errInvalidMultiSQL,
+ },
+ {
+ name: "update join is unsupported",
+ parseCtx: mustParseMultiExecutionPlanTestSQL(
+ t,
+ "UPDATE t_user u "+
+ "JOIN t_account a ON a.user_id = u.id "+
+ "SET u.status = 1 WHERE a.status = 0;"+
+ "UPDATE t_user u "+
+ "JOIN t_account a ON a.user_id = u.id "+
+ "SET u.status = 2 WHERE a.status = 1",
+ ),
+ dbType: types.DBTypeMySQL,
+ expectError: errUnsupportedMultiSQL,
+ },
+ {
+ name: "delete join is unsupported",
+ parseCtx: mustParseMultiExecutionPlanTestSQL(
+ t,
+ "DELETE u FROM t_user u "+
+ "JOIN t_account a ON a.user_id = u.id "+
+ "WHERE a.status = ?;"+
+ "DELETE u FROM t_user u "+
+ "JOIN t_account a ON a.user_id = u.id "+
+ "WHERE a.status = ?",
+ ),
+ dbType: types.DBTypeMySQL,
+ expectError: errUnsupportedMultiSQL,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ plan, err := buildMultiExecutionPlan(tt.parseCtx,
tt.dbType)
+
+ assert.Nil(t, plan)
+ assert.Error(t, err)
+ assert.ErrorIs(t, err, tt.expectError)
+ })
+ }
+}
+
+func mustParseMultiExecutionPlanTestSQL(t *testing.T, sourceQuery string)
*types.ParseContext {
+ t.Helper()
+
+ parseCtx, err := parser.DoParser(sourceQuery)
+ if !assert.NoError(t, err) {
+ t.FailNow()
+ }
+
+ return parseCtx
+}
diff --git a/pkg/datasource/sql/exec/at/multi_executor.go
b/pkg/datasource/sql/exec/at/multi_executor.go
index 8ce53a8e..0afe161d 100644
--- a/pkg/datasource/sql/exec/at/multi_executor.go
+++ b/pkg/datasource/sql/exec/at/multi_executor.go
@@ -20,12 +20,15 @@ package at
import (
"context"
"fmt"
+ "sync"
"seata.apache.org/seata-go/v2/pkg/datasource/sql/exec"
"seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
"seata.apache.org/seata-go/v2/pkg/util/log"
)
+var aggregateHookFallbackWarnOnce sync.Once
+
type multiExecutor struct {
baseExecutor
parserCtx *types.ParseContext
@@ -39,6 +42,54 @@ func NewMultiExecutor(parserCtx *types.ParseContext,
execContext *types.ExecCont
// ExecContext exec SQL, and generate before image and after image
func (m *multiExecutor) ExecContext(ctx context.Context, f
exec.CallbackWithNamedValue) (types.ExecResult, error) {
+ plan, err := buildMultiExecutionPlan(m.parserCtx, m.execContext.DBType)
+ if err != nil {
+ return nil, err
+ }
+
+ if plan.useAggregatePath {
+ if hasStatementSpecificHooks(plan) {
+ warnAggregateHookFallbackOnce()
+ } else {
+ return m.execAggregate(ctx, f, m.parserCtx)
+ }
+ }
+ return m.execSequential(ctx, f, plan)
+}
+
+func warnAggregateHookFallbackOnce() {
+ aggregateHookFallbackWarnOnce.Do(func() {
+ log.Warn(
+ "AT multi-SQL aggregate path skipped because
statement-specific hooks are registered globally; " +
+ "using sequential execution to preserve the
per-statement hook lifecycle",
+ )
+ })
+}
+
+func hasStatementSpecificHooks(plan *multiExecutionPlan) bool {
+ if plan == nil {
+ return false
+ }
+
+ for _, statementCtx := range plan.statements {
+ if statementCtx == nil {
+ continue
+ }
+
+ if len(hooksForSQLType(statementCtx.SQLType)) != 0 {
+ return true
+ }
+ }
+ return false
+}
+
+// execAggregate executes the optimized grouped UPDATE/DELETE path.
+//
+// It generates aggregate before images, executes the original multi-SQL once,
+// generates aggregate after images, validates them, and only then appends the
+// images to the transaction context. The callback result is returned
unchanged;
+// this executor does not aggregate RowsAffected or LastInsertId across
statements.
+func (m *multiExecutor) execAggregate(ctx context.Context, f
exec.CallbackWithNamedValue, parseCtx *types.ParseContext) (types.ExecResult,
error) {
if err := m.beforeHooks(ctx, m.execContext); err != nil {
return nil, err
}
@@ -47,30 +98,33 @@ func (m *multiExecutor) ExecContext(ctx context.Context, f
exec.CallbackWithName
m.afterHooks(ctx, m.execContext)
}()
- beforeImages, err := m.beforeImage(ctx, m.parserCtx)
+ beforeImages, err := m.beforeImage(ctx, parseCtx)
if err != nil {
return nil, err
}
- res, err := f(ctx, m.execContext.Query, m.execContext.NamedValues)
+ result, err := f(ctx, m.execContext.Query, m.execContext.NamedValues)
if err != nil {
return nil, err
}
- afterImages, err := m.afterImage(ctx, m.parserCtx, beforeImages)
+ afterImages, err := m.afterImage(ctx, parseCtx, beforeImages)
if err != nil {
return nil, err
}
- for _, beforeImage := range beforeImages {
- m.execContext.TxCtx.RoundImages.AppendBeofreImage(beforeImage)
+ if err := validateAggregateImages(parseCtx, beforeImages, afterImages);
err != nil {
+ return nil, err
}
- for _, afterImage := range afterImages {
- m.execContext.TxCtx.RoundImages.AppendAfterImage(afterImage)
+
+ for index := range beforeImages {
+
m.execContext.TxCtx.RoundImages.AppendBeofreImage(beforeImages[index])
+
m.execContext.TxCtx.RoundImages.AppendAfterImage(afterImages[index])
}
- return res, nil
+ return result, nil
}
+
func (m *multiExecutor) beforeImage(ctx context.Context, parseContext
*types.ParseContext) ([]*types.RecordImage, error) {
if len(parseContext.MultiStmt) == 0 {
return nil, nil
@@ -167,3 +221,34 @@ func (m *multiExecutor)
groupParsersByTableName(parseContext *types.ParseContext
return tableParsers, err
}
+
+func validateAggregateImages(parseCtx *types.ParseContext, beforeImages
[]*types.RecordImage, afterImages []*types.RecordImage) error {
+ if len(beforeImages) != len(afterImages) {
+ return fmt.Errorf("aggregate before/after image count mismatch:
"+"before=%d, after=%d", len(beforeImages), len(afterImages))
+ }
+
+ if parseCtx == nil || len(parseCtx.MultiStmt) == 0 ||
parseCtx.MultiStmt[0] == nil {
+ return fmt.Errorf("aggregate parse context contains no
statements")
+ }
+
+ executorType := parseCtx.MultiStmt[0].ExecutorType
+
+ for index := range beforeImages {
+ beforeImage := beforeImages[index]
+ afterImage := afterImages[index]
+
+ if beforeImage == nil || afterImage == nil {
+ return fmt.Errorf("aggregate image %d is nil", index)
+ }
+
+ if beforeImage.TableName != afterImage.TableName {
+ return fmt.Errorf("aggregate image %d table mismatch:
"+"before=%q, after=%q", index, beforeImage.TableName, afterImage.TableName)
+ }
+
+ if executorType == types.UpdateExecutor &&
len(beforeImage.Rows) != len(afterImage.Rows) {
+ return fmt.Errorf("aggregate update image %d row count
mismatch: "+"before=%d, after=%d", index, len(beforeImage.Rows),
len(afterImage.Rows))
+ }
+ }
+
+ return nil
+}
diff --git a/pkg/datasource/sql/exec/at/multi_executor_test.go
b/pkg/datasource/sql/exec/at/multi_executor_test.go
index e8cd6c33..70755af0 100644
--- a/pkg/datasource/sql/exec/at/multi_executor_test.go
+++ b/pkg/datasource/sql/exec/at/multi_executor_test.go
@@ -16,3 +16,70 @@
*/
package at
+
+import (
+ "context"
+ "database/sql/driver"
+ "testing"
+
+ "github.com/stretchr/testify/assert"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/exec"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/parser"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+)
+
+func TestMultiExecutorFallsBackToSequentialWhenStatementSpecificHookExists(t
*testing.T) {
+ sourceQuery := "UPDATE t_user SET name = 'user1' WHERE id = 1;" +
"UPDATE t_user SET age = 18 WHERE id = 2"
+
+ parseCtx, err := parser.DoParser(sourceQuery)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ plan, err := buildMultiExecutionPlan(parseCtx, types.DBTypeMySQL)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ assert.True(t, plan.useAggregatePath)
+
+ updateHook := &statementSpecificHookForTest{sqlType:
types.SQLTypeUpdate}
+
+ originalHooksForSQLType := hooksForSQLType
+ t.Cleanup(func() { hooksForSQLType = originalHooksForSQLType })
+
+ hooksForSQLType = func(sqlType types.SQLType) []exec.SQLHook {
+ if sqlType == types.SQLTypeUpdate {
+ return []exec.SQLHook{updateHook}
+ }
+ return nil
+ }
+
+ factoryRecorder := installSequentialFactoriesForTest(t)
+
+ execCtx := &types.ExecContext{
+ TxCtx: types.NewTxCtx(),
+ Query: sourceQuery,
+ ParseContext: parseCtx,
+ DBType: types.DBTypeMySQL,
+ }
+
+ multiExec := NewMultiExecutor(parseCtx, execCtx, nil)
+
+ callbackQueries := make([]string, 0, 2)
+
+ result, err := multiExec.ExecContext(context.Background(), func(ctx
context.Context, query string, args []driver.NamedValue) (types.ExecResult,
error) {
+ callbackQueries = append(callbackQueries, query)
+ return
types.NewResult(types.WithResult(driver.RowsAffected(1))), nil
+ },
+ )
+
+ if !assert.NoError(t, err) || !assert.NotNil(t, result) {
+ return
+ }
+
+ assert.Equal(t, []string{"UPDATE t_user SET name = 'user1' WHERE id =
1", "UPDATE t_user SET age = 18 WHERE id = 2"}, callbackQueries)
+ assert.Len(t, factoryRecorder.calls, 2)
+ assert.Equal(t, 2, updateHook.beforeCount)
+ assert.Equal(t, 2, updateHook.afterCount)
+}
diff --git a/pkg/datasource/sql/exec/at/multi_sequential_executor.go
b/pkg/datasource/sql/exec/at/multi_sequential_executor.go
new file mode 100644
index 00000000..c85f1c0b
--- /dev/null
+++ b/pkg/datasource/sql/exec/at/multi_sequential_executor.go
@@ -0,0 +1,262 @@
+/*
+ * 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 at
+
+import (
+ "bytes"
+ "context"
+ "database/sql/driver"
+ "fmt"
+ "strings"
+
+ "github.com/arana-db/parser/ast"
+ "github.com/arana-db/parser/format"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/exec"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/parser"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+)
+
+// execSequential executes validated statements in their original order.
+//
+// The first loop prepares all statement-local SQL and validates the total
argument count.
+// The second loop performs the actual business execution.
+// It intentionally returns the final successful statement's ExecResult to
+// match the aggregate path. RowsAffected and LastInsertId are not accumulated.
+func (m *multiExecutor) execSequential(ctx context.Context, f
exec.CallbackWithNamedValue, plan *multiExecutionPlan) (types.ExecResult,
error) {
+ if plan == nil {
+ return nil, fmt.Errorf("%w: execution plan is nil",
errInvalidMultiSQL)
+ }
+
+ if len(plan.statements) == 0 {
+ return nil, fmt.Errorf("%w: execution plan contains no
statements", errInvalidMultiSQL)
+ }
+
+ queries := make([]string, len(plan.statements))
+ argCounts := make([]int, len(plan.statements))
+ statementParseContexts := make([]*types.ParseContext,
len(plan.statements))
+ totalArgCount := 0
+
+ for index, statementCtx := range plan.statements {
+ stmtNode, err := getStatementNode(statementCtx)
+ if err != nil {
+ return nil, fmt.Errorf("get statement %d AST: %w",
index, err)
+ }
+
+ query, err := restoreStatementSQL(stmtNode)
+ if err != nil {
+ return nil, fmt.Errorf("restore statement %d SQL: %w",
index, err)
+ }
+
+ // Multi-statement parsing assigns parameter-marker orders
globally across
+ // all child statements. Reparse each restored child as
standalone SQL to
+ // rebase marker orders to the child-local argument slice
expected by the
+ // existing single-statement executors.
+ childParseCtx, err := parser.DoParser(query)
+ if err != nil {
+ return nil, fmt.Errorf("parse restored statement %d:
%w", index, err)
+ }
+
+ if childParseCtx.ExecutorType != statementCtx.ExecutorType {
+ return nil, fmt.Errorf(
+ "%w: statement %d changed executor type from %v
to %v after restoration",
+ errInvalidMultiSQL, index,
statementCtx.ExecutorType, childParseCtx.ExecutorType,
+ )
+ }
+
+ childStmtNode, err := getStatementNode(childParseCtx)
+ if err != nil {
+ return nil, fmt.Errorf("get restored statement %d AST:
%w", index, err)
+ }
+
+ argCount, err := countStatementParameters(childStmtNode)
+ if err != nil {
+ return nil, fmt.Errorf("count statement %d parameters:
%w", index, err)
+ }
+
+ queries[index] = query
+ argCounts[index] = argCount
+ statementParseContexts[index] = childParseCtx
+ totalArgCount += argCount
+ }
+
+ if totalArgCount != len(m.execContext.NamedValues) {
+ return nil, fmt.Errorf(
+ "%w: statements require %d arguments, but %d were
provided",
+ errInvalidMultiSQL, totalArgCount,
len(m.execContext.NamedValues),
+ )
+ }
+
+ if err := m.beforeHooks(ctx, m.execContext); err != nil {
+ return nil, err
+ }
+
+ defer func() {
+ m.afterHooks(ctx, m.execContext)
+ }()
+
+ var (
+ argOffset int
+ lastResult types.ExecResult
+ )
+
+ for index, childParseCtx := range statementParseContexts {
+ argCount := argCounts[index]
+ statementArgs :=
cloneNamedValues(m.execContext.NamedValues[argOffset : argOffset+argCount])
+
+ argOffset += argCount
+
+ childExecCtx := *m.execContext
+ childExecCtx.Query = queries[index]
+ childExecCtx.ParseContext = childParseCtx
+ childExecCtx.NamedValues = statementArgs
+ childExecCtx.Values = nil
+
+ childExecutor, err := newSequentialStatementExecutor(index,
childParseCtx, &childExecCtx)
+ if err != nil {
+ return nil, err
+ }
+
+ result, err := childExecutor.ExecContext(ctx, f)
+ if err != nil {
+ return nil, fmt.Errorf("execute statement %d: %w",
index, err)
+ }
+
+ if result == nil {
+ return nil, fmt.Errorf("%w: statement %d returned nil
result", errInvalidMultiSQL, index)
+ }
+
+ lastResult = result
+ }
+
+ return lastResult, nil
+}
+
+// newSequentialStatementExecutor reuses the existing single-statement
executors.
+//
+// Parent common and SQLTypeMulti hooks are owned by execSequential.
+// Each child executor receives only hooks registered for its own SQL type.
+func newSequentialStatementExecutor(index int, parseCtx *types.ParseContext,
execCtx *types.ExecContext) (executor, error) {
+ childHooks := hooksForSQLType(parseCtx.SQLType)
+ switch parseCtx.ExecutorType {
+ case types.InsertExecutor:
+ return newInsertExecutor(parseCtx, execCtx, childHooks), nil
+
+ case types.UpdateExecutor:
+ return newUpdateExecutor(parseCtx, execCtx, childHooks), nil
+
+ case types.DeleteExecutor:
+ return newDeleteExecutor(parseCtx, execCtx, childHooks), nil
+
+ default:
+ return nil, fmt.Errorf("%w: statement %d uses executor type %v",
+ errUnsupportedMultiSQL, index, parseCtx.ExecutorType,
+ )
+ }
+}
+
+// getStatementNode gets the concrete AST node from ParseContext.
+func getStatementNode(parseCtx *types.ParseContext) (ast.StmtNode, error) {
+ if parseCtx == nil {
+ return nil, fmt.Errorf("%w: statement parse context is nil",
errInvalidMultiSQL)
+ }
+
+ switch parseCtx.ExecutorType {
+ case types.InsertExecutor:
+ return parseCtx.InsertStmt, nil
+
+ case types.UpdateExecutor:
+ return parseCtx.UpdateStmt, nil
+
+ case types.DeleteExecutor:
+ return parseCtx.DeleteStmt, nil
+
+ default:
+ return nil, fmt.Errorf("%w: executor type %v",
errUnsupportedMultiSQL, parseCtx.ExecutorType)
+ }
+}
+
+// restoreStatementSQL obtains SQL that can be executed independently.
+//
+// Prefer the parser-preserved original statement text. If it is unavailable,
+// restore SQL from the AST.
+func restoreStatementSQL(stmt ast.StmtNode) (string, error) {
+ if stmt == nil {
+ return "", fmt.Errorf("%w: statement AST is nil",
errInvalidMultiSQL)
+ }
+
+ query := trimStatementSemicolon(stmt.OriginalText())
+ if query != "" {
+ return query, nil
+ }
+
+ var buffer bytes.Buffer
+ restoreCtx := format.NewRestoreCtx(format.DefaultRestoreFlags, &buffer)
+
+ if err := stmt.Restore(restoreCtx); err != nil {
+ return "", err
+ }
+
+ query = trimStatementSemicolon(buffer.String())
+ if query == "" {
+ return "", fmt.Errorf("%w: restored statement SQL is empty",
errInvalidMultiSQL)
+ }
+ return query, nil
+}
+
+func trimStatementSemicolon(query string) string {
+ query = strings.TrimSpace(query)
+
+ for strings.HasSuffix(query, ";") {
+ query = strings.TrimSpace(strings.TrimSuffix(query, ";"))
+ }
+ return query
+}
+
+type parameterCounter struct {
+ count int
+}
+
+func (c *parameterCounter) Enter(node ast.Node) (ast.Node, bool) {
+ if _, ok := node.(ast.ParamMarkerExpr); ok {
+ c.count++
+ }
+ return node, false
+}
+
+func (c *parameterCounter) Leave(node ast.Node) (ast.Node, bool) {
+ return node, true
+}
+
+func countStatementParameters(stmt ast.StmtNode) (int, error) {
+ counter := new(parameterCounter)
+
+ if _, ok := stmt.Accept(counter); !ok {
+ return 0, fmt.Errorf("%w: parameter traversal stopped
unexpectedly", errInvalidMultiSQL)
+ }
+
+ return counter.count, nil
+}
+
+func cloneNamedValues(values []driver.NamedValue) []driver.NamedValue {
+ cloned := make([]driver.NamedValue, len(values))
+ copy(cloned, values)
+ for index := range cloned {
+ cloned[index].Ordinal = index + 1
+ }
+ return cloned
+}
diff --git a/pkg/datasource/sql/exec/at/multi_sequential_executor_test.go
b/pkg/datasource/sql/exec/at/multi_sequential_executor_test.go
new file mode 100644
index 00000000..225aca1f
--- /dev/null
+++ b/pkg/datasource/sql/exec/at/multi_sequential_executor_test.go
@@ -0,0 +1,771 @@
+/*
+ * 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 at
+
+import (
+ "context"
+ "database/sql/driver"
+ "errors"
+ "fmt"
+ "testing"
+
+ "github.com/arana-db/parser/ast"
+ "github.com/arana-db/parser/test_driver"
+ "github.com/stretchr/testify/assert"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/exec"
+
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/parser"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+)
+
+func TestRestoreStatementSQLFromMultiSQL(t *testing.T) {
+ tests := []struct {
+ name string
+ sourceQuery string
+ expectQueries []string
+ expectExecutorTypes []types.ExecutorType
+ }{
+ {
+ name: "restore update and delete independently",
+ sourceQuery: "UPDATE t_user SET name = ? WHERE id = ?;"
+
+ "DELETE FROM t_user_log WHERE user_id = ?",
+ expectQueries: []string{
+ "UPDATE t_user SET name = ? WHERE id = ?",
+ "DELETE FROM t_user_log WHERE user_id = ?",
+ },
+ expectExecutorTypes:
[]types.ExecutorType{types.UpdateExecutor, types.DeleteExecutor},
+ },
+ {
+ name: "restore multiple inserts independently",
+ sourceQuery: "INSERT INTO t_user(id, name) VALUES (?,
?);" +
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ expectQueries: []string{
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ },
+ expectExecutorTypes:
[]types.ExecutorType{types.InsertExecutor, types.InsertExecutor},
+ },
+ {
+ name: "restore mixed statements in original order",
+ sourceQuery: "INSERT INTO t_user(id, name) VALUES (?,
?);" +
+ "UPDATE t_user SET name = ? WHERE id = ?;" +
+ "DELETE FROM t_user_log WHERE user_id = ?",
+ expectQueries: []string{
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ "UPDATE t_user SET name = ? WHERE id = ?",
+ "DELETE FROM t_user_log WHERE user_id = ?",
+ },
+ expectExecutorTypes:
[]types.ExecutorType{types.InsertExecutor, types.UpdateExecutor,
types.DeleteExecutor},
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ parseCtx, err := parser.DoParser(tt.sourceQuery)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ if !assert.Len(t, parseCtx.MultiStmt,
len(tt.expectQueries)) {
+ return
+ }
+
+ for index, statementCtx := range parseCtx.MultiStmt {
+ stmtNode, err := getStatementNode(statementCtx)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ query, err := restoreStatementSQL(stmtNode)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ assert.Equal(t, tt.expectQueries[index], query)
+ assert.NotContains(t, query, ";")
+
+ restoredParseCtx, err := parser.DoParser(query)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ assert.Empty(t, restoredParseCtx.MultiStmt)
+ assert.Equal(t, tt.expectExecutorTypes[index],
restoredParseCtx.ExecutorType)
+ assert.NotSame(t, statementCtx,
restoredParseCtx)
+ }
+ })
+ }
+}
+
+func TestCountStatementParameters(t *testing.T) {
+ tests := []struct {
+ name string
+ sourceQuery string
+ expectParameterCount int
+ }{
+ {
+ name: "insert with multiple rows",
+ sourceQuery: "INSERT INTO t_user(id, name)
VALUES (?, ?), (?, ?)",
+ expectParameterCount: 4,
+ },
+ {
+ name: "update parameters in set where
and limit",
+ sourceQuery: "UPDATE t_user SET name = ?, age
= ? " + "WHERE id = ? LIMIT ?",
+ expectParameterCount: 4,
+ },
+ {
+ name: "question mark in string literal
is not parameter",
+ sourceQuery: "UPDATE t_user SET description =
'ready?' " + "WHERE id = ?",
+ expectParameterCount: 1,
+ },
+ {
+ name: "delete without parameters",
+ sourceQuery: "DELETE FROM t_user WHERE id = 1",
+ expectParameterCount: 0,
+ },
+ {
+ name: "delete with multiple parameters",
+ sourceQuery: "DELETE FROM t_user " + "WHERE
status = ? AND created_at < ?",
+ expectParameterCount: 2,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ parseCtx, err := parser.DoParser(tt.sourceQuery)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ stmtNode, err := getStatementNode(parseCtx)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ parameterCount, err :=
countStatementParameters(stmtNode)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ assert.Equal(t, tt.expectParameterCount, parameterCount)
+ })
+ }
+}
+
+func TestCloneNamedValues(t *testing.T) {
+ tests := []struct {
+ name string
+ sourceNamedValues []driver.NamedValue
+ expectNamedValues []driver.NamedValue
+ }{
+ {
+ name: "rebase single named value",
+ sourceNamedValues: []driver.NamedValue{{Name:
"user_id", Ordinal: 3, Value: int64(10)}},
+ expectNamedValues: []driver.NamedValue{{Name:
"user_id", Ordinal: 1, Value: int64(10)}},
+ },
+ {
+ name: "rebase multiple named values",
+ sourceNamedValues: []driver.NamedValue{{Ordinal: 4,
Value: "user"}, {Ordinal: 5, Value: int64(18)}, {Ordinal: 6, Value: int64(10)}},
+ expectNamedValues: []driver.NamedValue{{Ordinal: 1,
Value: "user"}, {Ordinal: 2, Value: int64(18)}, {Ordinal: 3, Value: int64(10)}},
+ },
+ {
+ name: "clone empty named values",
+ sourceNamedValues: []driver.NamedValue{},
+ expectNamedValues: []driver.NamedValue{},
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ originalNamedValues := make([]driver.NamedValue,
len(tt.sourceNamedValues))
+ copy(originalNamedValues, tt.sourceNamedValues)
+
+ cloned := cloneNamedValues(tt.sourceNamedValues)
+ assert.Equal(t, tt.expectNamedValues, cloned)
+ assert.Equal(t, originalNamedValues,
tt.sourceNamedValues)
+
+ if len(cloned) == 0 {
+ return
+ }
+
+ cloned[0].Ordinal = 100
+ cloned[0].Value = "changed"
+
+ assert.Equal(t, originalNamedValues,
tt.sourceNamedValues)
+ })
+ }
+}
+
+func TestReparseStatementResetsParameterOrders(t *testing.T) {
+ sourceQuery := "UPDATE t_user SET name = ? WHERE id = ?;" + "DELETE
FROM t_user_log WHERE user_id = ?"
+
+ parseCtx, err := parser.DoParser(sourceQuery)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ if !assert.Len(t, parseCtx.MultiStmt, 2) {
+ return
+ }
+
+ originalUpdateNode, err := getStatementNode(parseCtx.MultiStmt[0])
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ originalDeleteNode, err := getStatementNode(parseCtx.MultiStmt[1])
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ assert.Equal(t, []int{0, 1}, collectParameterOrdersForTest(t,
originalUpdateNode))
+ assert.Equal(t, []int{2}, collectParameterOrdersForTest(t,
originalDeleteNode))
+
+ deleteQuery, err := restoreStatementSQL(originalDeleteNode)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ reparsedDeleteCtx, err := parser.DoParser(deleteQuery)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ reparsedDeleteNode, err := getStatementNode(reparsedDeleteCtx)
+ if !assert.NoError(t, err) {
+ return
+ }
+
+ assert.Equal(t, []int{0}, collectParameterOrdersForTest(t,
reparsedDeleteNode))
+
+ globalNamedValues := []driver.NamedValue{{Ordinal: 1, Value: "user"},
{Ordinal: 2, Value: int64(10)}, {Ordinal: 3, Value: int64(10)}}
+ deleteNamedValues := cloneNamedValues(globalNamedValues[2:])
+
+ assert.Equal(t, []driver.NamedValue{{Ordinal: 1, Value: int64(10)}},
deleteNamedValues)
+ assert.Equal(t, []int{2}, collectParameterOrdersForTest(t,
originalDeleteNode))
+}
+
+type parameterOrderCollectorForTest struct {
+ orders []int
+}
+
+func (c *parameterOrderCollectorForTest) Enter(node ast.Node) (ast.Node, bool)
{
+ if marker, ok := node.(*test_driver.ParamMarkerExpr); ok {
+ c.orders = append(c.orders, marker.Order)
+ }
+
+ return node, false
+}
+
+func (c *parameterOrderCollectorForTest) Leave(node ast.Node) (ast.Node, bool)
{
+ return node, true
+}
+
+func collectParameterOrdersForTest(t *testing.T, stmt ast.StmtNode) []int {
+ t.Helper()
+
+ if !assert.NotNil(t, stmt) {
+ t.FailNow()
+ }
+
+ collector := new(parameterOrderCollectorForTest)
+
+ if _, ok := stmt.Accept(collector); !assert.True(t, ok) {
+ t.FailNow()
+ }
+
+ return collector.orders
+}
+
+func TestExecSequentialPreservesOrderAndArguments(t *testing.T) {
+ tests := []struct {
+ name string
+ sourceQuery string
+ namedValues []driver.NamedValue
+ expectQueries []string
+ expectedNamedValues [][]driver.NamedValue
+ expectExecutorTypes []types.ExecutorType
+ }{
+ {
+ name: "multiple inserts execute in original order",
+ sourceQuery: "INSERT INTO t_user(id, name) VALUES (?,
?);" +
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ namedValues: sequentialNamedValuesForTest(int64(1),
"user1", int64(2), "user2"),
+ expectQueries: []string{
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ },
+ expectedNamedValues: [][]driver.NamedValue{
+ sequentialNamedValuesForTest(int64(1), "user1"),
+ sequentialNamedValuesForTest(int64(2), "user2"),
+ },
+ expectExecutorTypes:
[]types.ExecutorType{types.InsertExecutor, types.InsertExecutor},
+ },
+ {
+ name: "mixed DML executes in original order",
+ sourceQuery: "INSERT INTO t_user(id, name) VALUES (?,
?);" +
+ "UPDATE t_user SET name = ? WHERE id = ?;" +
+ "DELETE FROM t_user WHERE id = ?",
+ namedValues: sequentialNamedValuesForTest(int64(1),
"user1", "Updated user1", int64(1), int64(1)),
+ expectQueries: []string{
+ "INSERT INTO t_user(id, name) VALUES (?, ?)",
+ "UPDATE t_user SET name = ? WHERE id = ?",
+ "DELETE FROM t_user WHERE id = ?",
+ },
+ expectedNamedValues: [][]driver.NamedValue{
+ sequentialNamedValuesForTest(int64(1), "user1"),
+ sequentialNamedValuesForTest("Updated user1",
int64(1)),
+ sequentialNamedValuesForTest(int64(1)),
+ },
+ expectExecutorTypes:
[]types.ExecutorType{types.InsertExecutor, types.UpdateExecutor,
types.DeleteExecutor},
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ factoryRecorder := installSequentialFactoriesForTest(t)
+ var events []string
+ hook := &sequentialHookForTest{events: &events}
+
+ multiExec := newMultiExecutorForSequentialTest(t,
tt.sourceQuery, tt.namedValues, []exec.SQLHook{hook})
+ executions := make([]sequentialExecutionForTest, 0,
len(tt.expectQueries))
+
+ result, err :=
multiExec.ExecContext(context.Background(), func(ctx context.Context, query
string, args []driver.NamedValue) (types.ExecResult, error) {
+ executions = append(executions,
sequentialExecutionForTest{
+ query: query,
+ namedValues:
cloneNamedValuesForSequentialTest(args),
+ })
+
+ events = append(events, "execute:"+query)
+ return
types.NewResult(types.WithResult(driver.RowsAffected(len(executions)))), nil
+ })
+
+ if !assert.NoError(t, err) || !assert.NotNil(t, result)
{
+ return
+ }
+
+ assertSequentialFactoryCallsForTest(t,
factoryRecorder.calls, tt.expectExecutorTypes, tt.expectQueries,
tt.expectedNamedValues)
+
+ expectedExecutions :=
make([]sequentialExecutionForTest, len(tt.expectQueries))
+ for index := range tt.expectQueries {
+ expectedExecutions[index] =
sequentialExecutionForTest{
+ query: tt.expectQueries[index],
+ namedValues:
tt.expectedNamedValues[index],
+ }
+ }
+
+ assert.Equal(t, expectedExecutions, executions)
+
+ // Hooks belong to the complete multi-SQL operation,
rather than to each child statement.
+ assert.Equal(t, 1, hook.beforeCount)
+ assert.Equal(t, 1, hook.afterCount)
+ assert.Same(t, multiExec.execContext,
hook.beforeExecCtx)
+ assert.Same(t, multiExec.execContext, hook.afterExecCtx)
+
+ expectedEvents := []string{"before"}
+ for _, query := range tt.expectQueries {
+ expectedEvents = append(expectedEvents,
"execute:"+query)
+ }
+
+ expectedEvents = append(expectedEvents, "after")
+ assert.Equal(t, expectedEvents, events)
+
+ rowsAffected, err := result.GetResult().RowsAffected()
+ if assert.NoError(t, err) {
+ assert.Equal(t, int64(len(tt.expectQueries)),
rowsAffected)
+ }
+
+ })
+ }
+}
+
+func TestExecSequentialReturnsFinalStatementResult(t *testing.T) {
+ sourceQuery := "INSERT INTO t_user(id, name) VALUES (?, ?);" +
+ "INSERT INTO t_user(id, name) VALUES (?, ?)"
+ namedValues := sequentialNamedValuesForTest(int64(1), "user1",
int64(2), "user2")
+ results := []*mockExecResult{
+ {lastInsertID: 101, rowsAffected: 2},
+ {lastInsertID: 202, rowsAffected: 5},
+ }
+
+ installSequentialFactoriesForTest(t)
+ multiExec := newMultiExecutorForSequentialTest(t, sourceQuery,
namedValues, nil)
+ executionCount := 0
+
+ result, err := multiExec.ExecContext(context.Background(),
func(context.Context, string, []driver.NamedValue) (types.ExecResult, error) {
+ result := results[executionCount]
+ executionCount++
+ return result, nil
+ })
+
+ if !assert.NoError(t, err) || !assert.Same(t, results[1], result) {
+ return
+ }
+ assert.Equal(t, len(results), executionCount)
+
+ rowsAffected, err := result.GetResult().RowsAffected()
+ if assert.NoError(t, err) {
+ assert.Equal(t, int64(5), rowsAffected)
+ }
+
+ lastInsertID, err := result.GetResult().LastInsertId()
+ if assert.NoError(t, err) {
+ assert.Equal(t, int64(202), lastInsertID)
+ }
+}
+
+func TestExecSequentialStopsOnMiddleFailure(t *testing.T) {
+ sourceQuery := "INSERT INTO t_user(id, name) VALUES (?,?);" +
+ "INSERT INTO t_user(id, name) VALUES (?,?);" +
+ "INSERT INTO t_user(id, name) VALUES (?,?)"
+
+ nameValues := sequentialNamedValuesForTest(int64(1), "user1", int64(2),
"user2", int64(3), "user3")
+ factoryRecorder := installSequentialFactoriesForTest(t)
+
+ var events []string
+ hook := &sequentialHookForTest{
+ events: &events,
+ }
+
+ multiExec := newMultiExecutorForSequentialTest(t, sourceQuery,
nameValues, []exec.SQLHook{hook})
+ expectedError := errors.New("second statement failed")
+ executions := make([]sequentialExecutionForTest, 0, 2)
+
+ result, err := multiExec.ExecContext(context.Background(), func(ctx
context.Context, query string, args []driver.NamedValue) (types.ExecResult,
error) {
+ executions = append(executions, sequentialExecutionForTest{
+ query: query,
+ namedValues: cloneNamedValuesForSequentialTest(args),
+ })
+
+ events = append(events, "execute:"+query)
+ if len(executions) == 2 {
+ return nil, expectedError
+ }
+
+ return
types.NewResult(types.WithResult(driver.RowsAffected(1))), nil
+ },
+ )
+
+ assert.Nil(t, result)
+ if !assert.ErrorIs(t, err, expectedError) {
+ return
+ }
+
+ assert.Contains(t, err.Error(), "execute statement 1")
+
+ assert.Len(t, factoryRecorder.calls, 2)
+ assert.Len(t, executions, 2)
+
+ expectedQuery := "INSERT INTO t_user(id, name) VALUES (?,?)"
+
+ assert.Equal(t, expectedQuery, executions[0].query)
+ assert.Equal(t, expectedQuery, executions[1].query)
+
+ assert.Equal(t, sequentialNamedValuesForTest(int64(1), "user1"),
executions[0].namedValues)
+ assert.Equal(t, sequentialNamedValuesForTest(int64(2), "user2"),
executions[1].namedValues)
+
+ assert.Equal(t, 1, hook.beforeCount)
+ assert.Equal(t, 1, hook.afterCount)
+
+ assert.Equal(t, []string{"before", "execute:" + expectedQuery,
"execute:" + expectedQuery, "after"}, events)
+}
+
+func TestExecSequentialRejectsArgumentCountMismatchBeforeSideEffects(t
*testing.T) {
+ sourceQuery := "INSERT INTO t_user(id,name) VALUES (?,?);" + "INSERT
INTO t_user(id,name) VALUES (?,?)"
+
+ tests := []struct {
+ name string
+ namedValues []driver.NamedValue
+ }{
+ {
+ name: "too few arguments",
+ namedValues: sequentialNamedValuesForTest(int64(1),
"user1", int64(2)),
+ },
+ {
+ name: "too many arguments",
+ namedValues: sequentialNamedValuesForTest(int64(1),
"user1", int64(2), "user2", "unexpected"),
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ factoryRecorder := installSequentialFactoriesForTest(t)
+ hook := new(sequentialHookForTest)
+
+ multiExec := newMultiExecutorForSequentialTest(t,
sourceQuery, tt.namedValues, []exec.SQLHook{hook})
+ callbackCount := 0
+ result, err :=
multiExec.ExecContext(context.Background(), func(ctx context.Context, query
string, args []driver.NamedValue) (types.ExecResult, error) {
+ callbackCount++
+ return
types.NewResult(types.WithResult(driver.RowsAffected(1))), nil
+ })
+
+ assert.Nil(t, result)
+ if !assert.ErrorIs(t, err, errInvalidMultiSQL) {
+ return
+ }
+ assert.Contains(t, err.Error(), "statements require 4
arguments")
+ assert.Contains(t, err.Error(), fmt.Sprintf("but %d
were provided", len(tt.namedValues)))
+
+ assert.Empty(t, factoryRecorder.calls)
+ assert.Zero(t, callbackCount)
+ assert.Zero(t, hook.beforeCount)
+ assert.Zero(t, hook.afterCount)
+ })
+ }
+}
+
+type sequentialExecutionForTest struct {
+ query string
+ namedValues []driver.NamedValue
+}
+
+type sequentialFactoryCallForTest struct {
+ factoryExecutorType types.ExecutorType
+ parseExecutorType types.ExecutorType
+
+ query string
+ namedValues []driver.NamedValue
+
+ childHookCount int
+ parseContextMatch bool
+ isSingleStatement bool
+}
+
+type sequentialFactoryRecorderForTest struct {
+ calls []sequentialFactoryCallForTest
+}
+
+func (r *sequentialFactoryRecorderForTest) build(factoryExecutorType
types.ExecutorType, parseCtx *types.ParseContext,
+ execCtx *types.ExecContext, hooks []exec.SQLHook,
+) executor {
+ r.calls = append(r.calls,
+ sequentialFactoryCallForTest{
+ factoryExecutorType: factoryExecutorType,
+ parseExecutorType: parseCtx.ExecutorType,
+ query: execCtx.Query,
+ namedValues:
cloneNamedValuesForSequentialTest(execCtx.NamedValues),
+ childHookCount: len(hooks),
+ parseContextMatch: parseCtx == execCtx.ParseContext,
+ isSingleStatement: len(parseCtx.MultiStmt) == 0,
+ },
+ )
+
+ return &sequentialCallbackExecutorForTest{baseExecutor:
baseExecutor{hooks: append([]exec.SQLHook(nil), hooks...)}, execCtx: execCtx}
+}
+
+type sequentialCallbackExecutorForTest struct {
+ baseExecutor
+ execCtx *types.ExecContext
+}
+
+func (e *sequentialCallbackExecutorForTest) ExecContext(ctx context.Context, f
exec.CallbackWithNamedValue) (types.ExecResult, error) {
+ if err := e.beforeHooks(ctx, e.execCtx); err != nil {
+ return nil, err
+ }
+ defer func() {
+ e.afterHooks(ctx, e.execCtx)
+ }()
+ return f(ctx, e.execCtx.Query, e.execCtx.NamedValues)
+}
+
+type sequentialHookForTest struct {
+ beforeCount int
+ afterCount int
+
+ beforeExecCtx *types.ExecContext
+ afterExecCtx *types.ExecContext
+
+ events *[]string
+}
+
+func (h *sequentialHookForTest) Type() types.SQLType {
+ return types.SQLTypeMulti
+}
+
+func (h *sequentialHookForTest) Before(ctx context.Context, execCtx
*types.ExecContext) error {
+ h.beforeCount++
+ h.beforeExecCtx = execCtx
+ if h.events != nil {
+ *h.events = append(*h.events, "before")
+ }
+ return nil
+}
+
+func (h *sequentialHookForTest) After(ctx context.Context, execCtx
*types.ExecContext) error {
+ h.afterCount++
+ h.afterExecCtx = execCtx
+ if h.events != nil {
+ *h.events = append(*h.events, "after")
+ }
+ return nil
+}
+
+func installSequentialFactoriesForTest(t *testing.T)
*sequentialFactoryRecorderForTest {
+ t.Helper()
+
+ originalInsertExecutor := newInsertExecutor
+ originalUpdateExecutor := newUpdateExecutor
+ originalDeleteExecutor := newDeleteExecutor
+
+ t.Cleanup(func() {
+ newInsertExecutor = originalInsertExecutor
+ newUpdateExecutor = originalUpdateExecutor
+ newDeleteExecutor = originalDeleteExecutor
+ })
+
+ recorder := new(sequentialFactoryRecorderForTest)
+
+ newInsertExecutor = func(parseCtx *types.ParseContext, execCtx
*types.ExecContext, hooks []exec.SQLHook) executor {
+ return recorder.build(types.InsertExecutor, parseCtx, execCtx,
hooks)
+ }
+
+ newUpdateExecutor = func(parseCtx *types.ParseContext, execCtx
*types.ExecContext, hooks []exec.SQLHook) executor {
+ return recorder.build(types.UpdateExecutor, parseCtx, execCtx,
hooks)
+ }
+
+ newDeleteExecutor = func(parseCtx *types.ParseContext, execCtx
*types.ExecContext, hooks []exec.SQLHook) executor {
+ return recorder.build(types.DeleteExecutor, parseCtx, execCtx,
hooks)
+ }
+
+ return recorder
+}
+
+func newMultiExecutorForSequentialTest(t *testing.T, sourceQuery string,
namedValues []driver.NamedValue, hooks []exec.SQLHook) *multiExecutor {
+ t.Helper()
+
+ parseCtx, err := parser.DoParser(sourceQuery)
+ if !assert.NoError(t, err) {
+ t.FailNow()
+ }
+
+ execCtx := &types.ExecContext{
+ TxCtx: types.NewTxCtx(),
+ Query: sourceQuery,
+ ParseContext: parseCtx,
+ NamedValues: namedValues,
+ DBType: types.DBTypeMySQL,
+ }
+
+ builtExecutor := NewMultiExecutor(parseCtx, execCtx, hooks)
+
+ multiExec, ok := builtExecutor.(*multiExecutor)
+ if !assert.True(t, ok) {
+ t.FailNow()
+ }
+
+ return multiExec
+}
+
+func sequentialNamedValuesForTest(values ...driver.Value) []driver.NamedValue {
+ namedValues := make([]driver.NamedValue, len(values))
+
+ for index, value := range values {
+ namedValues[index] = driver.NamedValue{
+ Ordinal: index + 1,
+ Value: value,
+ }
+ }
+
+ return namedValues
+}
+
+func cloneNamedValuesForSequentialTest(values []driver.NamedValue)
[]driver.NamedValue {
+ cloned := make([]driver.NamedValue, len(values))
+ copy(cloned, values)
+ return cloned
+}
+
+func assertSequentialFactoryCallsForTest(t *testing.T, calls
[]sequentialFactoryCallForTest,
+ expectExecutorTypes []types.ExecutorType, expectQueries []string,
expectNamedValues [][]driver.NamedValue) {
+ t.Helper()
+
+ if !assert.Len(t, calls, len(expectQueries)) {
+ return
+ }
+
+ for index, call := range calls {
+ assert.Equal(t, expectExecutorTypes[index],
call.factoryExecutorType)
+ assert.Equal(t, expectExecutorTypes[index],
call.parseExecutorType)
+ assert.Equal(t, expectQueries[index], call.query)
+ assert.Equal(t, expectNamedValues[index], call.namedValues)
+
+ assert.True(t, call.parseContextMatch)
+ assert.True(t, call.isSingleStatement)
+
+ assert.Zero(t, call.childHookCount)
+ }
+}
+
+func TestExecSequentialRunsInsertSpecificHook(t *testing.T) {
+ expectedErr := errors.New("insert blocked")
+ insertHook := &statementSpecificHookForTest{sqlType:
types.SQLTypeInsert, beforeErr: expectedErr}
+
+ originalHooksForSQLType := hooksForSQLType
+ t.Cleanup(func() {
+ hooksForSQLType = originalHooksForSQLType
+ })
+
+ hooksForSQLType = func(sqlType types.SQLType) []exec.SQLHook {
+ if sqlType == types.SQLTypeInsert {
+ return []exec.SQLHook{insertHook}
+ }
+ return nil
+ }
+
+ sourceQuery := "INSERT INTO forbidden_table(id) VALUES (?);" + "INSERT
INTO forbidden_table(id) VALUES (?)"
+
+ multiExec := newMultiExecutorForSequentialTest(t, sourceQuery,
sequentialNamedValuesForTest(int64(1), int64(2)), nil)
+ callbackCount := 0
+ result, err := multiExec.ExecContext(context.Background(),
+ func(ctx context.Context, query string, args
[]driver.NamedValue) (types.ExecResult, error) {
+ callbackCount++
+ return
types.NewResult(types.WithResult(driver.RowsAffected(1))), nil
+ },
+ )
+
+ assert.Nil(t, result)
+ assert.ErrorIs(t, err, expectedErr)
+ assert.Zero(t, callbackCount)
+ assert.Equal(t, 1, insertHook.beforeCount)
+ assert.Zero(t, insertHook.afterCount)
+}
+
+type statementSpecificHookForTest struct {
+ sqlType types.SQLType
+ beforeErr error
+ beforeCount int
+ afterCount int
+ beforeExecCtx *types.ExecContext
+ afterExecCtx *types.ExecContext
+}
+
+func (h *statementSpecificHookForTest) Type() types.SQLType {
+ return h.sqlType
+}
+
+func (h *statementSpecificHookForTest) Before(ctx context.Context, execCtx
*types.ExecContext) error {
+ h.beforeCount++
+ h.beforeExecCtx = execCtx
+ return h.beforeErr
+}
+
+func (h *statementSpecificHookForTest) After(ctx context.Context, execCtx
*types.ExecContext) error {
+ h.afterCount++
+ h.afterExecCtx = execCtx
+ return nil
+}
diff --git a/pkg/datasource/sql/exec/at/multi_update_excutor.go
b/pkg/datasource/sql/exec/at/multi_update_excutor.go
index 5e784793..9c26a9dc 100644
--- a/pkg/datasource/sql/exec/at/multi_update_excutor.go
+++ b/pkg/datasource/sql/exec/at/multi_update_excutor.go
@@ -45,7 +45,6 @@ type multiUpdateExecutor struct {
execContext *types.ExecContext
}
-var rows driver.Rows
var comma = ","
// NewMultiUpdateExecutor get new multi update executor
@@ -87,7 +86,7 @@ func (u *multiUpdateExecutor) ExecContext(ctx
context.Context, f exec.CallbackWi
}
for i, afterImage := range afterImages {
- beforeImage := afterImages[i]
+ beforeImage := beforeImages[i]
if len(beforeImage.Rows) != len(afterImage.Rows) {
return nil, errors.New("Before image size is not
equaled to after image size, probably because you updated the primary keys.")
}
@@ -117,15 +116,19 @@ func (u *multiUpdateExecutor) beforeImage(ctx
context.Context) ([]*types.RecordI
}
rows, err := u.rowsPrepare(ctx, selectSQL, selectArgs)
+ if err != nil {
+ return nil, err
+ }
+
defer func() {
- if err := rows.Close(); err != nil {
- log.Errorf("rows close fail, err:%v", err)
+ if rows == nil {
return
}
+
+ if closeErr := rows.Close(); closeErr != nil {
+ log.Errorf("rows close fail,err: %v", closeErr)
+ }
}()
- if err != nil {
- return nil, err
- }
image, err := u.buildRecordImages(rows, metaData, types.SQLTypeUpdate,
types.DBTypeMySQL)
if err != nil {
@@ -148,6 +151,9 @@ func (u *multiUpdateExecutor) afterImage(ctx
context.Context, beforeImages []*ty
return nil, errors.New("empty beforeImages")
}
beforeImage := beforeImages[0]
+ if beforeImage == nil {
+ return nil, errors.New("aggregate update before image is nil")
+ }
tableName :=
u.parserCtx.MultiStmt[0].UpdateStmt.TableRefs.TableRefs.Left.(*ast.TableSource).Source.(*ast.TableName).Name.O
metaData, err :=
datasource.GetTableCache(types.DBTypeMySQL).GetTableMeta(ctx,
u.execContext.DBName, tableName)
@@ -155,19 +161,30 @@ func (u *multiUpdateExecutor) afterImage(ctx
context.Context, beforeImages []*ty
return nil, err
}
+ // No row matched the aggregate UPDATE predicates.
+ //
+ // Do not generate an after-image SELECT with an empty primary-key list.
+ // Return one empty image so that before/after image counts remain
equal.
+ if len(beforeImage.Rows) == 0 {
+ return []*types.RecordImage{types.NewEmptyRecordImage(metaData,
u.parserCtx.SQLType)}, nil
+ }
+
// use
selectSQL, selectArgs := u.buildAfterImageSQL(beforeImage, *metaData)
- rows, err = u.rowsPrepare(ctx, selectSQL, selectArgs)
+ rows, err := u.rowsPrepare(ctx, selectSQL, selectArgs)
+ if err != nil {
+ return nil, err
+ }
defer func() {
- if err := rows.Close(); err != nil {
- log.Errorf("rows close fail, err:%v", err)
+ if rows == nil {
return
}
+
+ if closeErr := rows.Close(); closeErr != nil {
+ log.Errorf("rows close fail,err: %v", closeErr)
+ }
}()
- if err != nil {
- return nil, err
- }
image, err := u.buildRecordImages(rows, metaData, types.SQLTypeUpdate,
types.DBTypeMySQL)
if err != nil {
@@ -179,25 +196,12 @@ func (u *multiUpdateExecutor) afterImage(ctx
context.Context, beforeImages []*ty
}
func (u *multiUpdateExecutor) rowsPrepare(ctx context.Context, selectSQL
string, selectArgs []driver.NamedValue) (driver.Rows, error) {
- var queryer driver.Queryer
-
- queryerContext, ok := u.execContext.Conn.(driver.QueryerContext)
- if !ok {
- queryer, ok = u.execContext.Conn.(driver.Queryer)
- }
- if ok {
- var err error
- rows, err = util.CtxDriverQuery(ctx, queryerContext, queryer,
selectSQL, selectArgs)
-
- if err != nil {
- log.Errorf("ctx driver query: %+v", err)
- return nil, err
- }
- } else {
- log.Errorf("target conn should been driver.QueryerContext or
driver.Queryer")
- return nil, errors.New("invalid conn")
+ rowsi, err := util.CtxDriverQueryWithPrepareFallback(ctx,
u.execContext.Conn, selectSQL, selectArgs)
+ if err != nil {
+ log.Errorf("aggregate update image query failed,err: %+v", err)
+ return nil, err
}
- return rows, nil
+ return rowsi, nil
}
// buildAfterImageSQL build the SQL to query after image data
diff --git a/pkg/datasource/sql/exec/at/update_join_executor.go
b/pkg/datasource/sql/exec/at/update_join_executor.go
index dad8a507..ca388c2c 100644
--- a/pkg/datasource/sql/exec/at/update_join_executor.go
+++ b/pkg/datasource/sql/exec/at/update_join_executor.go
@@ -133,7 +133,7 @@ func (u *updateJoinExecutor) beforeImage(ctx
context.Context) ([]*types.RecordIm
image, err = u.buildRecordImages(rowsi, metaData,
types.SQLTypeUpdate, types.DBTypeMySQL)
}
if rowsi != nil {
- if rowerr := rows.Close(); rowerr != nil {
+ if rowerr := rowsi.Close(); rowerr != nil {
log.Errorf("rows close fail, err:%v", rowerr)
return nil, rowerr
}
diff --git a/pkg/datasource/sql/exec/hook.go b/pkg/datasource/sql/exec/hook.go
index 0a81db52..129e85ab 100644
--- a/pkg/datasource/sql/exec/hook.go
+++ b/pkg/datasource/sql/exec/hook.go
@@ -66,3 +66,13 @@ type SQLHook interface {
Before(ctx context.Context, execCtx *types.ExecContext) error
After(ctx context.Context, execCtx *types.ExecContext) error
}
+
+// HooksForSQLType returns a copy of hooks registered for the given SQL type.
+//
+// Common hooks are intentionally excluded. MultiExecutor owns the common and
+// SQLTypeMulti hook lifecycle, while each sequential child runs only its
+// statement-specific hooks.
+func HooksForSQLType(sqlType types.SQLType) []SQLHook {
+ hooks := hookSolts[sqlType]
+ return append([]SQLHook(nil), hooks...)
+}
diff --git a/pkg/datasource/sql/util/ctxutil.go
b/pkg/datasource/sql/util/ctxutil.go
index fa29440b..e5beae68 100644
--- a/pkg/datasource/sql/util/ctxutil.go
+++ b/pkg/datasource/sql/util/ctxutil.go
@@ -121,3 +121,92 @@ func namedValueToValue(named []driver.NamedValue)
([]driver.Value, error) {
}
return dargs, nil
}
+
+type rowsWithStmt struct {
+ driver.Rows
+ stmt driver.Stmt
+}
+
+func (r *rowsWithStmt) Close() error {
+ var rowsErr error
+ if r.Rows != nil {
+ rowsErr = r.Rows.Close()
+ }
+
+ var stmtErr error
+ if r.stmt != nil {
+ stmtErr = r.stmt.Close()
+ }
+
+ return errors.Join(rowsErr, stmtErr)
+}
+
+// CtxDriverExecWithPrepareFallback first tries the connection-level Exec path.
+// If the driver returns driver.ErrSkip, it prepares and executes the statement
+// directly on the underlying driver connection.
+func CtxDriverExecWithPrepareFallback(ctx context.Context, conn driver.Conn,
query string, args []driver.NamedValue) (driver.Result, error) {
+ var execerContext driver.ExecerContext
+ if execer, ok := conn.(driver.ExecerContext); ok {
+ execerContext = execer
+ }
+
+ var execer driver.Execer
+ if e, ok := conn.(driver.Execer); ok {
+ execer = e
+ }
+
+ if execerContext != nil || execer != nil {
+ result, err := ctxDriverExec(ctx, execerContext, execer, query,
args)
+ if err == nil {
+ return result, nil
+ }
+
+ if !errors.Is(err, driver.ErrSkip) {
+ return nil, err
+ }
+ }
+
+ stmt, err := ctxDriverPrepare(ctx, conn, query)
+ if err != nil {
+ return nil, err
+ }
+ defer stmt.Close()
+
+ return ctxDriverStmtExec(ctx, stmt, args)
+}
+
+func CtxDriverQueryWithPrepareFallback(ctx context.Context, conn driver.Conn,
query string, args []driver.NamedValue) (driver.Rows, error) {
+ var queryerContext driver.QueryerContext
+ if queryer, ok := conn.(driver.QueryerContext); ok {
+ queryerContext = queryer
+ }
+
+ var queryer driver.Queryer
+ if q, ok := conn.(driver.Queryer); ok {
+ queryer = q
+ }
+
+ if queryerContext != nil || queryer != nil {
+ rows, err := CtxDriverQuery(ctx, queryerContext, queryer,
query, args)
+ if err == nil {
+ return rows, nil
+ }
+
+ if !errors.Is(err, driver.ErrSkip) {
+ return nil, err
+ }
+ }
+
+ stmt, err := ctxDriverPrepare(ctx, conn, query)
+ if err != nil {
+ return nil, err
+ }
+
+ rows, err := ctxDriverStmtQuery(ctx, stmt, args)
+ if err != nil {
+ _ = stmt.Close()
+ return nil, err
+ }
+
+ return &rowsWithStmt{Rows: rows, stmt: stmt}, nil
+}
diff --git a/pkg/datasource/sql/util/ctxutil_test.go
b/pkg/datasource/sql/util/ctxutil_test.go
index 51bef2f6..415bd266 100644
--- a/pkg/datasource/sql/util/ctxutil_test.go
+++ b/pkg/datasource/sql/util/ctxutil_test.go
@@ -23,7 +23,9 @@ import (
"errors"
"testing"
+ "github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/mock"
)
// Mock implementations for testing
@@ -503,3 +505,165 @@ func TestNamedValueToValue_Empty(t *testing.T) {
assert.NoError(t, err)
assert.Empty(t, result)
}
+
+func TestCtxDriverExecWithPrepareFallback(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ ctx := context.Background()
+ query := "INSERT INTO t_user(id, name) VALUES (?, ?)"
+ args := []driver.NamedValue{{Ordinal: 1, Value: int64(1)}, {Ordinal: 2,
Value: "user1"}}
+
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+ targetStmt := mock.NewMockTestDriverStmt(ctrl)
+
+ targetConn.EXPECT().ExecContext(ctx, query, args).Return(nil,
driver.ErrSkip)
+ targetConn.EXPECT().PrepareContext(ctx, query).Return(targetStmt, nil)
+ targetStmt.EXPECT().ExecContext(ctx,
args).Return(driver.RowsAffected(1), nil)
+ targetStmt.EXPECT().Close().Return(nil)
+
+ result, err := CtxDriverExecWithPrepareFallback(ctx, targetConn, query,
args)
+
+ if !assert.NoError(t, err) || !assert.NotNil(t, result) {
+ return
+ }
+
+ affected, err := result.RowsAffected()
+ assert.NoError(t, err)
+ assert.Equal(t, int64(1), affected)
+}
+
+func TestCtxDriverQueryWithPrepareFallback(t *testing.T) {
+ ctx := context.Background()
+ query := "SELECT id, name FROM t_user WHERE id IN (?, ?)"
+ args := []driver.NamedValue{
+ {
+ Ordinal: 1,
+ Value: int64(1),
+ },
+ {
+ Ordinal: 2,
+ Value: int64(2),
+ },
+ }
+
+ t.Run("direct query success does not prepare", func(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+ targetRows := mock.NewMockTestDriverRows(ctrl)
+
+ targetConn.EXPECT().QueryContext(ctx, query,
args).Times(1).Return(targetRows, nil)
+ targetRows.EXPECT().Close().Times(1).Return(nil)
+
+ rows, err := CtxDriverQueryWithPrepareFallback(ctx, targetConn,
query, args)
+
+ if !assert.NoError(t, err) || !assert.NotNil(t, rows) {
+ return
+ }
+
+ assert.Same(t, targetRows, rows)
+ assert.NoError(t, rows.Close())
+ })
+
+ t.Run("ErrSkip falls back to prepared query", func(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+ targetStmt := mock.NewMockTestDriverStmt(ctrl)
+ targetRows := mock.NewMockTestDriverRows(ctrl)
+
+ targetConn.EXPECT().QueryContext(ctx, query,
args).Times(1).Return(nil, driver.ErrSkip)
+ targetConn.EXPECT().PrepareContext(ctx,
query).Times(1).Return(targetStmt, nil)
+ targetStmt.EXPECT().QueryContext(ctx,
args).Times(1).Return(targetRows, nil)
+ targetRows.EXPECT().Close().Times(1).Return(nil)
+ targetStmt.EXPECT().Close().Times(1).Return(nil)
+
+ rows, err := CtxDriverQueryWithPrepareFallback(ctx, targetConn,
query, args)
+
+ if !assert.NoError(t, err) || !assert.NotNil(t, rows) {
+ return
+ }
+
+ assert.NotSame(t, targetRows, rows)
+ assert.NoError(t, rows.Close())
+ })
+
+ t.Run("non ErrSkip error does not prepare", func(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+ expectedErr := errors.New("direct query failed")
+
+ targetConn.EXPECT().QueryContext(ctx, query,
args).Times(1).Return(nil, expectedErr)
+ rows, err := CtxDriverQueryWithPrepareFallback(ctx, targetConn,
query, args)
+
+ assert.Nil(t, rows)
+ assert.ErrorIs(t, err, expectedErr)
+ })
+
+ t.Run("prepare error is returned", func(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+ expectedErr := errors.New("prepare query failed")
+
+ targetConn.EXPECT().QueryContext(ctx, query,
args).Times(1).Return(nil, driver.ErrSkip)
+ targetConn.EXPECT().PrepareContext(ctx,
query).Times(1).Return(nil, expectedErr)
+ rows, err := CtxDriverQueryWithPrepareFallback(ctx, targetConn,
query, args)
+
+ assert.Nil(t, rows)
+ assert.ErrorIs(t, err, expectedErr)
+ })
+
+ t.Run("prepared query error closes statement", func(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+ targetStmt := mock.NewMockTestDriverStmt(ctrl)
+ expectedErr := errors.New("prepared query failed")
+
+ targetConn.EXPECT().QueryContext(ctx, query,
args).Times(1).Return(nil, driver.ErrSkip)
+ targetConn.EXPECT().PrepareContext(ctx,
query).Times(1).Return(targetStmt, nil)
+ targetStmt.EXPECT().QueryContext(ctx,
args).Times(1).Return(nil, expectedErr)
+ targetStmt.EXPECT().Close().Times(1).Return(nil)
+
+ rows, err := CtxDriverQueryWithPrepareFallback(ctx, targetConn,
query, args)
+ assert.Nil(t, rows)
+ assert.ErrorIs(t, err, expectedErr)
+ })
+
+ t.Run("close preserves rows and statement errors", func(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ targetConn := mock.NewMockTestDriverConn(ctrl)
+ targetStmt := mock.NewMockTestDriverStmt(ctrl)
+ targetRows := mock.NewMockTestDriverRows(ctrl)
+
+ rowsCloseErr := errors.New("rows close failed")
+ stmtCloseErr := errors.New("statement close failed")
+
+ targetConn.EXPECT().QueryContext(ctx, query,
args).Times(1).Return(nil, driver.ErrSkip)
+ targetConn.EXPECT().PrepareContext(ctx,
query).Times(1).Return(targetStmt, nil)
+ targetStmt.EXPECT().QueryContext(ctx,
args).Times(1).Return(targetRows, nil)
+ targetRows.EXPECT().Close().Times(1).Return(rowsCloseErr)
+ targetStmt.EXPECT().Close().Times(1).Return(stmtCloseErr)
+
+ rows, err := CtxDriverQueryWithPrepareFallback(ctx, targetConn,
query, args)
+
+ if !assert.NoError(t, err) || !assert.NotNil(t, rows) {
+ return
+ }
+
+ closeErr := rows.Close()
+
+ assert.ErrorIs(t, closeErr, rowsCloseErr)
+ assert.ErrorIs(t, closeErr, stmtCloseErr)
+ })
+}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]