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 3bcd201f feat: support handling of composite primary keys and related
tests (#1133)
3bcd201f is described below
commit 3bcd201fd293614314e018bbc6077046ff0c77e3
Author: tachibana22 <[email protected]>
AuthorDate: Thu Aug 27 14:38:34 2026 +0800
feat: support handling of composite primary keys and related tests (#1133)
* feat: support handling of composite primary keys and related tests
* feat(util): add nil guards and case-insensitive matching for pk list
parsing
* feat(exec): optimize micro-allocations and support composite pk in insert
executor
* feat(undo): align composite autoincrement and array merge for insert undo
builder
* fix(undo): enforce sentinel token leak prevention and enrich semantic
error context
* feat(executor): implement length-prefixed encoding for composite primary
keys in rowListToMap
* feat(test): add test for GetOrderedPkList with shuffled composite primary
keys
* fix(tests): improve formatting and ensure newline at end of file in test
files
---
pkg/datasource/sql/datasource/mysql/trigger.go | 2 +-
.../sql/datasource/mysql/trigger_test.go | 24 ++++
pkg/datasource/sql/exec/at/insert_executor.go | 16 +--
pkg/datasource/sql/exec/at/insert_executor_test.go | 55 +++++++++
pkg/datasource/sql/types/meta_test.go | 27 +++++
.../undo/builder/mysql_insert_undo_log_builder.go | 11 +-
.../builder/mysql_insert_undo_log_builder_test.go | 26 ++++
.../undo/executor/mysql_undo_insert_executor.go | 13 +-
.../executor/mysql_undo_insert_executor_test.go | 77 ++++++++++++
pkg/datasource/sql/undo/executor/utils.go | 37 +++---
pkg/datasource/sql/undo/executor/utils_test.go | 131 +++++++++++++++++++--
pkg/datasource/sql/util/escape.go | 15 ++-
pkg/datasource/sql/util/escape_test.go | 60 ++++++++++
13 files changed, 442 insertions(+), 52 deletions(-)
diff --git a/pkg/datasource/sql/datasource/mysql/trigger.go
b/pkg/datasource/sql/datasource/mysql/trigger.go
index a14189c5..519393c4 100644
--- a/pkg/datasource/sql/datasource/mysql/trigger.go
+++ b/pkg/datasource/sql/datasource/mysql/trigger.go
@@ -178,7 +178,7 @@ func (m *mysqlTrigger) getIndexes(ctx context.Context,
dbName string, tableName
tableName = util.DelEscape(tableName, types.DBTypeMySQL)
result := make([]types.IndexMeta, 0)
- indexMetaSql := "SELECT `INDEX_NAME`, `COLUMN_NAME`, `NON_UNIQUE` FROM
`INFORMATION_SCHEMA`.`STATISTICS` WHERE `TABLE_SCHEMA` = ? AND `TABLE_NAME` = ?"
+ indexMetaSql := "SELECT `INDEX_NAME`, `COLUMN_NAME`, `NON_UNIQUE` FROM
`INFORMATION_SCHEMA`.`STATISTICS` WHERE `TABLE_SCHEMA` = ? AND `TABLE_NAME` = ?
ORDER BY `SEQ_IN_INDEX` ASC"
stmt, err := conn.PrepareContext(ctx, indexMetaSql)
if err != nil {
return nil, err
diff --git a/pkg/datasource/sql/datasource/mysql/trigger_test.go
b/pkg/datasource/sql/datasource/mysql/trigger_test.go
index ab9037c6..0f49042b 100644
--- a/pkg/datasource/sql/datasource/mysql/trigger_test.go
+++ b/pkg/datasource/sql/datasource/mysql/trigger_test.go
@@ -552,6 +552,30 @@ func Test_mysqlTrigger_getIndexes(t *testing.T) {
assert.Equal(t, types.IndexTypePrimaryKey,
indexes[0].IType)
},
},
+ {
+ name: "success_composite_primary_key",
+ setupMock: func() {
+ rows := sqlmock.NewRows([]string{"INDEX_NAME",
"COLUMN_NAME", "NON_UNIQUE"}).
+ AddRow("PRIMARY", "tenant_id", 0).
+ AddRow("PRIMARY", "id", 0)
+
+ mock.ExpectPrepare("SELECT (.+) FROM
`INFORMATION_SCHEMA`.`STATISTICS`").
+ ExpectQuery().
+ WithArgs("testdb", "users").
+ WillReturnRows(rows)
+ },
+ expectError: false,
+ expectedCount: 2,
+ validateIndex: func(t *testing.T, indexes
[]types.IndexMeta) {
+ assert.Equal(t, types.IndexTypePrimaryKey,
indexes[0].IType)
+ assert.Equal(t, "PRIMARY", indexes[0].Name)
+ assert.Equal(t, "tenant_id",
indexes[0].ColumnName)
+
+ assert.Equal(t, types.IndexTypePrimaryKey,
indexes[1].IType)
+ assert.Equal(t, "PRIMARY", indexes[1].Name)
+ assert.Equal(t, "id", indexes[1].ColumnName)
+ },
+ },
}
for _, tt := range tests {
diff --git a/pkg/datasource/sql/exec/at/insert_executor.go
b/pkg/datasource/sql/exec/at/insert_executor.go
index c1b9e791..200a59e8 100644
--- a/pkg/datasource/sql/exec/at/insert_executor.go
+++ b/pkg/datasource/sql/exec/at/insert_executor.go
@@ -345,22 +345,22 @@ func (i *insertExecutor) buildAfterImageSQL(ctx
context.Context) (string, []driv
if len(dataTypeMap) != len(pkColumnNameList) {
return "", nil, fmt.Errorf("PK columnName size don't equal PK
DataType size")
}
- var pkRowImages []types.RowImage
rowSize := len(pkValuesMap[pkColumnNameList[0]])
+ pkRowImages := make([]types.RowImage, 0, rowSize)
for i := 0; i < rowSize; i++ {
+ columns := make([]types.ColumnImage, 0, len(pkColumnNameList))
for _, name := range pkColumnNameList {
tmpKey := name
tmpArray := pkValuesMap[tmpKey]
- pkRowImages = append(pkRowImages, types.RowImage{
- Columns: []types.ColumnImage{{
- KeyType: types.IndexTypePrimaryKey,
- ColumnName: tmpKey,
- ColumnType:
jdbcTypeForDatabaseType(dbType, dataTypeMap[tmpKey]),
- Value: tmpArray[i],
- }},
+ columns = append(columns, types.ColumnImage{
+ KeyType: types.IndexTypePrimaryKey,
+ ColumnName: tmpKey,
+ ColumnType: jdbcTypeForDatabaseType(dbType,
dataTypeMap[tmpKey]),
+ Value: tmpArray[i],
})
}
+ pkRowImages = append(pkRowImages, types.RowImage{Columns:
columns})
}
// build check sql
sb := strings.Builder{}
diff --git a/pkg/datasource/sql/exec/at/insert_executor_test.go
b/pkg/datasource/sql/exec/at/insert_executor_test.go
index b3a2463e..f5eb103f 100644
--- a/pkg/datasource/sql/exec/at/insert_executor_test.go
+++ b/pkg/datasource/sql/exec/at/insert_executor_test.go
@@ -326,6 +326,29 @@ func TestBuildSelectSQLByInsert(t *testing.T) {
expectQuery: "SELECT id, tenant_id, name FROM user
WHERE (`id`,`tenant_id`) IN ((?,?),(?,?)) ",
expectQueryArgs: []driver.Value{int64(19), int64(100),
int64(21), int64(101)},
},
+ {
+ name: "test-composite-pk-allocation",
+ query: "insert into user(tenant_id, id, name) values
('tenantA', 100, 'Tony'), ('tenantB', 101, 'Tom')",
+ metaData: types.TableMeta{
+ ColumnNames: []string{"tenant_id", "id",
"name"},
+ Indexs: map[string]types.IndexMeta{
+ "PRIMARY": {
+ IType:
types.IndexTypePrimaryKey,
+ Columns: []types.ColumnMeta{
+ {ColumnName:
"tenant_id", DatabaseType: types.GetSqlDataType("VARCHAR")},
+ {ColumnName: "id",
DatabaseType: types.GetSqlDataType("BIGINT")},
+ },
+ },
+ },
+ Columns: map[string]types.ColumnMeta{
+ "tenant_id": {ColumnName: "tenant_id"},
+ "id": {ColumnName: "id"},
+ "name": {ColumnName: "name"},
+ },
+ },
+ expectQuery: "SELECT tenant_id, id, name FROM user
WHERE (`tenant_id`,`id`) IN ((?,?),(?,?)) ",
+ expectQueryArgs: []driver.Value{"tenantA", int64(100),
"tenantB", int64(101)},
+ },
}
for _, test := range tests {
@@ -1118,6 +1141,38 @@ func TestMySQLInsertUndoLogBuilder_autoGeneratePks(t
*testing.T) {
}
}
+func TestCanAutoGeneratePKs_CompositePK(t *testing.T) {
+ tests := []struct {
+ name string
+ pkMetaMap map[string]types.ColumnMeta
+ want bool
+ }{
+ {
+ name: "composite primary key with one autoincrement
column",
+ pkMetaMap: map[string]types.ColumnMeta{
+ "tenant_id": {ColumnName: "tenant_id",
Autoincrement: false},
+ "id": {ColumnName: "id", Autoincrement:
true},
+ },
+ want: true,
+ },
+ {
+ name: "composite primary key without any autoincrement
column",
+ pkMetaMap: map[string]types.ColumnMeta{
+ "group_id": {ColumnName: "group_id",
Autoincrement: false},
+ "user_id": {ColumnName: "user_id",
Autoincrement: false},
+ },
+ want: false,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got := canAutoGeneratePKs(tt.pkMetaMap)
+ assert.Equal(t, tt.want, got)
+ })
+ }
+}
+
type autoIncrementStepConn struct {
value driver.Value
}
diff --git a/pkg/datasource/sql/types/meta_test.go
b/pkg/datasource/sql/types/meta_test.go
index d6ec45cf..bdc7beac 100644
--- a/pkg/datasource/sql/types/meta_test.go
+++ b/pkg/datasource/sql/types/meta_test.go
@@ -84,6 +84,33 @@ func TestTableMeta_GetPrimaryKeyTypeStrMap(t *testing.T) {
}}, want: map[string]string{
"id": "BIGINT",
}},
+ {
+ name: "test-composite-pk",
+ fields: fields{
+ TableName: "test_composite",
+ Indexs: map[string]IndexMeta{
+ "PRIMARY": {
+ Name: "PRIMARY",
+ ColumnName: "id",
+ IType: IndexTypePrimaryKey,
+ Columns: []ColumnMeta{
+ {
+ ColumnName:
"tenant_id",
+
DatabaseTypeString: "VARCHAR",
+ },
+ {
+ ColumnName:
"id",
+
DatabaseTypeString: "BIGINT",
+ },
+ },
+ },
+ },
+ },
+ want: map[string]string{
+ "tenant_id": "VARCHAR",
+ "id": "BIGINT",
+ },
+ },
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
diff --git a/pkg/datasource/sql/undo/builder/mysql_insert_undo_log_builder.go
b/pkg/datasource/sql/undo/builder/mysql_insert_undo_log_builder.go
index bb88654b..a1f227fe 100644
--- a/pkg/datasource/sql/undo/builder/mysql_insert_undo_log_builder.go
+++ b/pkg/datasource/sql/undo/builder/mysql_insert_undo_log_builder.go
@@ -113,9 +113,9 @@ func (u *MySQLInsertUndoLogBuilder) buildAfterImageSQL(ctx
context.Context, exec
if len(dataTypeMap) != len(pkColumnNameList) {
return "", nil, fmt.Errorf("PK columnName size don't equal PK
DataType size")
}
- var pkRowImages []types.RowImage
rowSize := len(pkValuesMap[pkColumnNameList[0]])
+ pkRowImages := make([]types.RowImage, 0, rowSize*len(pkColumnNameList))
for i := 0; i < rowSize; i++ {
for _, name := range pkColumnNameList {
tmpKey := name
@@ -445,11 +445,10 @@ func (u *MySQLInsertUndoLogBuilder)
getPkValuesByAuto(execCtx *types.ExecContext
}
func canAutoIncrement(pkMetaMap map[string]types.ColumnMeta) bool {
- if len(pkMetaMap) != 1 {
- return false
- }
for _, meta := range pkMetaMap {
- return meta.Autoincrement
+ if meta.Autoincrement {
+ return true
+ }
}
return false
}
@@ -504,7 +503,7 @@ func pkValuesMapMerge(dest *map[string][]interface{}, src
map[string][]interface
for k, v := range src {
tmpK := k
tmpV := v
- (*dest)[tmpK] = append((*dest)[tmpK], tmpV)
+ (*dest)[tmpK] = append((*dest)[tmpK], tmpV...)
}
}
diff --git
a/pkg/datasource/sql/undo/builder/mysql_insert_undo_log_builder_test.go
b/pkg/datasource/sql/undo/builder/mysql_insert_undo_log_builder_test.go
index b995b886..518f1b64 100644
--- a/pkg/datasource/sql/undo/builder/mysql_insert_undo_log_builder_test.go
+++ b/pkg/datasource/sql/undo/builder/mysql_insert_undo_log_builder_test.go
@@ -310,6 +310,32 @@ func TestBuildSelectSQLByInsert(t *testing.T) {
orExpectQuery: "SELECT * FROM user WHERE
(`name`,`id`) IN ((?,?),(?,?)) ",
orExpectQueryArgs: []driver.Value{"Tony", int64(19),
"Tom", int64(20)},
},
+ {
+ name: "test-composite-autoincrement-shadow-path",
+ query: "insert into user(tenant_id, name) values
('tenantX', 'Jack')",
+ metaDataMap: map[string]types.TableMeta{
+ "user": {
+ ColumnNames: []string{"tenant_id",
"id", "name"},
+ Indexs: map[string]types.IndexMeta{
+ "PRIMARY": {
+ IType:
types.IndexTypePrimaryKey,
+ Columns:
[]types.ColumnMeta{
+ {ColumnName:
"tenant_id", DatabaseType: types.GetSqlDataType("VARCHAR"), Autoincrement:
false},
+ {ColumnName:
"id", DatabaseType: types.GetSqlDataType("BIGINT"), Autoincrement: true},
+ },
+ },
+ },
+ Columns: map[string]types.ColumnMeta{
+ "tenant_id": {ColumnName:
"tenant_id", Autoincrement: false},
+ "id": {ColumnName: "id",
Autoincrement: true},
+ "name": {ColumnName:
"name", Autoincrement: false},
+ },
+ },
+ },
+ mockInsertResult: NewMockInsertResult(500, 1),
+ expectQuery: "SELECT * FROM user WHERE
(`tenant_id`,`id`) IN ((?,?)) ",
+ expectQueryArgs: []driver.Value{"tenantX", int64(500)},
+ },
}
for _, test := range tests {
diff --git a/pkg/datasource/sql/undo/executor/mysql_undo_insert_executor.go
b/pkg/datasource/sql/undo/executor/mysql_undo_insert_executor.go
index a3cf0b5d..9ff456fa 100644
--- a/pkg/datasource/sql/undo/executor/mysql_undo_insert_executor.go
+++ b/pkg/datasource/sql/undo/executor/mysql_undo_insert_executor.go
@@ -58,12 +58,15 @@ func (m *mySQLUndoInsertExecutor) ExecuteOn(ctx
context.Context, dbType types.DB
defer stmt.Close()
afterImage := m.sqlUndoLog.AfterImage
for _, row := range afterImage.Rows {
- pkValueList := make([]interface{}, 0)
+ pkList, err := util.GetOrderedPkList(afterImage, row, dbType)
+ if err != nil {
+ return fmt.Errorf("UNDO-INSERT-CONTEXT-ERROR [Op:
ExecuteOn, Table: %s]: failed to parse ordered primary keys from record image:
%w", m.sqlUndoLog.TableName, err)
+ }
+
+ pkValueList := make([]interface{}, 0, len(pkList))
- for _, col := range row.Columns {
- if col.KeyType == types.PrimaryKey.Number() {
- pkValueList = append(pkValueList, col.Value)
- }
+ for _, col := range pkList {
+ pkValueList = append(pkValueList, col.Value)
}
if _, err = stmt.Exec(pkValueList...); err != nil {
diff --git
a/pkg/datasource/sql/undo/executor/mysql_undo_insert_executor_test.go
b/pkg/datasource/sql/undo/executor/mysql_undo_insert_executor_test.go
index 0f9be1ff..5b2801d5 100644
--- a/pkg/datasource/sql/undo/executor/mysql_undo_insert_executor_test.go
+++ b/pkg/datasource/sql/undo/executor/mysql_undo_insert_executor_test.go
@@ -20,6 +20,7 @@ package executor
import (
"context"
"database/sql"
+ "fmt"
"testing"
"github.com/DATA-DOG/go-sqlmock"
@@ -307,6 +308,27 @@ func TestMySQLUndoInsertExecutor_ExecuteOn(t *testing.T) {
WillReturnError(assert.AnError)
},
},
+ {
+ name: "execute with ordered pk list error",
+ afterImage: &types.RecordImage{
+ TableName: "test_table",
+ TableMeta: &types.TableMeta{TableName:
"test_table"},
+ Rows: []types.RowImage{
+ {
+ Columns: []types.ColumnImage{
+ {ColumnName: "id",
KeyType: types.PrimaryKey.Number(), Value: 1},
+ },
+ },
+ },
+ },
+ expectError: true,
+ setupMock: func(mock sqlmock.Sqlmock) {
+ patches :=
gomonkey.ApplyFunc(util.GetOrderedPkList, func(image *types.RecordImage, row
types.RowImage, dbType types.DBType) ([]types.ColumnImage, error) {
+ return nil, fmt.Errorf("mock ordered pk
error")
+ })
+ t.Cleanup(func() { patches.Reset() })
+ },
+ },
}
for _, tt := range tests {
@@ -375,3 +397,58 @@ func TestMySQLUndoInsertExecutor_ExecuteOn(t *testing.T) {
})
}
}
+func TestMySQLUndoInsertExecutor_BuildUndoSQL_CompositePK(t *testing.T) {
+ afterImage := &types.RecordImage{
+ TableName: "test_table",
+ TableMeta: &types.TableMeta{
+ TableName: "test_table",
+ Columns: map[string]types.ColumnMeta{
+ "tenant_id": {ColumnName: "tenant_id"},
+ "id": {ColumnName: "id"},
+ },
+ Indexs: map[string]types.IndexMeta{
+ "PRIMARY": {
+ IType: types.IndexTypePrimaryKey,
+ Columns: []types.ColumnMeta{
+ {ColumnName: "tenant_id"},
+ {ColumnName: "id"},
+ },
+ },
+ },
+ },
+ Rows: []types.RowImage{
+ {
+ Columns: []types.ColumnImage{
+ {ColumnName: "tenant_id", KeyType:
types.IndexTypePrimaryKey, Value: "tenant_1"},
+ {ColumnName: "id", KeyType:
types.IndexTypePrimaryKey, Value: 100},
+ },
+ },
+ },
+ }
+
+ sqlUndoLog := undo.SQLUndoLog{
+ TableName: "test_table",
+ AfterImage: afterImage,
+ }
+
+ patches := gomonkey.ApplyFunc(util.GetOrderedPkList, func(image
*types.RecordImage, row types.RowImage, dbType types.DBType)
([]types.ColumnImage, error) {
+ return []types.ColumnImage{
+ {ColumnName: "tenant_id", Value: "tenant_1"},
+ {ColumnName: "id", Value: 100},
+ }, nil
+ })
+ defer patches.Reset()
+
+ patches.ApplyFunc(util.BuildWhereConditionByPKs, func(pkNameList
[]string, dbType types.DBType) string {
+ return "`" + pkNameList[0] + "` = ? AND `" + pkNameList[1] + "`
= ?"
+ })
+
+ executor := &mySQLUndoInsertExecutor{
+ sqlUndoLog: sqlUndoLog,
+ }
+
+ gotSQL, err := executor.buildUndoSQL(types.DBTypeMySQL)
+
+ assert.NoError(t, err)
+ assert.Equal(t, "DELETE FROM test_table WHERE `tenant_id` = ? AND `id`
= ? ", gotSQL)
+}
diff --git a/pkg/datasource/sql/undo/executor/utils.go
b/pkg/datasource/sql/undo/executor/utils.go
index 7c8b767c..60060d4e 100644
--- a/pkg/datasource/sql/undo/executor/utils.go
+++ b/pkg/datasource/sql/undo/executor/utils.go
@@ -68,27 +68,36 @@ func compareRows(tableMeta types.TableMeta, oldRows
[]types.RowImage, newRows []
}
func rowListToMap(rows []types.RowImage, primaryKeyList []string)
map[string]map[string]interface{} {
- rowMap := make(map[string]map[string]interface{}, 0)
- for _, row := range rows {
- fieldMap := make(map[string]interface{}, 0)
- var rowKey string
- var firstUnderline bool
+ rowMap := make(map[string]map[string]interface{}, len(rows))
+ for rowIndex, row := range rows {
+ fieldMap := make(map[string]interface{}, len(row.Columns))
+ pkValues := make(map[string]interface{}, len(primaryKeyList))
for _, column := range row.Columns {
cleanName := util.DelEscape(column.ColumnName,
types.DBTypeMySQL)
- for i, key := range primaryKeyList {
- if cleanName == key {
- if firstUnderline && i > 0 {
- rowKey += "_##$$_"
- }
- // todo make value more accurate
- rowKey = fmt.Sprintf("%v%v", rowKey,
column.GetActualValue())
- firstUnderline = true
+ for _, key := range primaryKeyList {
+ if strings.EqualFold(cleanName, key) {
+ pkValues[key] = column.GetActualValue()
}
}
fieldMap[strings.ToUpper(cleanName)] = column.Value
}
- rowMap[rowKey] = fieldMap
+
+ var sb strings.Builder
+ for i, key := range primaryKeyList {
+ if i > 0 {
+ sb.WriteString(",")
+ }
+ val, ok := pkValues[key]
+ var valStr string
+ if !ok || val == nil {
+ valStr =
fmt.Sprintf("__SENTINEL_MISSING_PK_%s_ROW_%d__", key, rowIndex)
+ } else {
+ valStr = fmt.Sprintf("%v", val)
+ }
+ sb.WriteString(fmt.Sprintf("%d:%s", len(valStr),
valStr))
+ }
+ rowMap[sb.String()] = fieldMap
}
return rowMap
}
diff --git a/pkg/datasource/sql/undo/executor/utils_test.go
b/pkg/datasource/sql/undo/executor/utils_test.go
index dff7b953..c977c205 100644
--- a/pkg/datasource/sql/undo/executor/utils_test.go
+++ b/pkg/datasource/sql/undo/executor/utils_test.go
@@ -362,20 +362,21 @@ func TestRowListToMap(t *testing.T) {
if len(tt.rows) > 0 {
for _, row := range tt.rows {
- // Verify that each row can be found in
the map by constructing expected key
- var expectedKey string
- var firstUnderline bool
- for _, column := range row.Columns {
- for i, key := range
tt.primaryKeyList {
- if column.ColumnName ==
key {
- if
firstUnderline && i > 0 {
-
expectedKey += "_##$$_"
- }
- expectedKey +=
fmt.Sprintf("%v", column.GetActualValue())
- firstUnderline
= true
+ // Verify that each row can be found in
the map by constructing expected key using Length-Prefixed Encoding
+ var sb strings.Builder
+ for i, key := range tt.primaryKeyList {
+ if i > 0 {
+ sb.WriteString(",")
+ }
+ for _, column := range
row.Columns {
+ if
strings.EqualFold(column.ColumnName, key) {
+ valStr :=
fmt.Sprintf("%v", column.GetActualValue())
+
sb.WriteString(fmt.Sprintf("%d:%s", len(valStr), valStr))
+ break
}
}
}
+ expectedKey := sb.String()
rowData, exists := result[expectedKey]
assert.True(t, exists, "Row should
exist in map with key: %s", expectedKey)
@@ -537,8 +538,8 @@ func TestRowListToMap_EscapedColumnNames(t *testing.T) {
result := rowListToMap(rows, primaryKeyList)
assert.Len(t, result, 2)
- // Verify rows can be found by their PK values
- row1, exists := result["1"]
+ // Verify rows can be found by their PK values (encoded as len:val ->
1:1)
+ row1, exists := result["1:1"]
assert.True(t, exists, "Row with PK=1 should exist")
if exists {
// After fix, fieldMap key uses cleaned (unescaped) uppercase
column name
@@ -569,3 +570,107 @@ func TestBuildPKParams_EscapedColumnNames(t *testing.T) {
assert.Len(t, result, 2)
assert.Equal(t, []interface{}{1, 2}, result)
}
+
+func TestRowListToMap_CompositePK_ColumnOrderIndependent(t *testing.T) {
+ primaryKeyList := []string{"tenant_id", "id"}
+
+ tests := []struct {
+ name string
+ rows []types.RowImage
+ want map[string]bool
+ }{
+ {
+ name: "physical columns order: tenant_id then id",
+ rows: []types.RowImage{
+ {
+ Columns: []types.ColumnImage{
+ {ColumnName: "tenant_id",
Value: "tenant123"},
+ {ColumnName: "id", Value: 456},
+ {ColumnName: "name", Value:
"test_a"},
+ },
+ },
+ },
+ want: map[string]bool{"9:tenant123,3:456": true},
+ },
+ {
+ name: "physical columns order: id then tenant_id
(shuffled)",
+ rows: []types.RowImage{
+ {
+ Columns: []types.ColumnImage{
+ {ColumnName: "id", Value: 456},
+ {ColumnName: "name", Value:
"test_b"},
+ {ColumnName: "tenant_id",
Value: "tenant123"},
+ },
+ },
+ },
+ want: map[string]bool{"9:tenant123,3:456": true},
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ gotMap := rowListToMap(tt.rows, primaryKeyList)
+
+ assert.Len(t, gotMap, 1)
+ for gotKey := range gotMap {
+ assert.True(t, tt.want[gotKey], "generated
rowKey %s not matching expected order", gotKey)
+ }
+ })
+ }
+}
+
+func TestRowListToMap_SentinelMissingPK(t *testing.T) {
+ primaryKeyList := []string{"tenant_id", "id"}
+ rows := []types.RowImage{
+ {
+ Columns: []types.ColumnImage{
+ {ColumnName: "id", Value: 789},
+ {ColumnName: "name", Value: "test_sentinel"},
+ },
+ },
+ }
+
+ gotMap := rowListToMap(rows, primaryKeyList)
+ assert.Len(t, gotMap, 1)
+
+ for gotKey := range gotMap {
+ assert.Contains(t, gotKey,
"__SENTINEL_MISSING_PK_tenant_id_ROW_0__")
+ assert.Contains(t, gotKey, "789")
+ }
+}
+
+func TestRowListToMap_CollisionPrevention(t *testing.T) {
+ primaryKeyList := []string{"pk1", "pk2"}
+
+ rowsA := []types.RowImage{
+ {
+ Columns: []types.ColumnImage{
+ {ColumnName: "pk1", Value: "a_##$$_b"},
+ {ColumnName: "pk2", Value: "c"},
+ },
+ },
+ }
+ rowsB := []types.RowImage{
+ {
+ Columns: []types.ColumnImage{
+ {ColumnName: "pk1", Value: "a"},
+ {ColumnName: "pk2", Value: "b_##$$_c"},
+ },
+ },
+ }
+
+ mapA := rowListToMap(rowsA, primaryKeyList)
+ mapB := rowListToMap(rowsB, primaryKeyList)
+
+ var keyA, keyB string
+ for k := range mapA {
+ keyA = k
+ }
+ for k := range mapB {
+ keyB = k
+ }
+
+ assert.Equal(t, "8:a_##$$_b,1:c", keyA)
+ assert.Equal(t, "1:a,8:b_##$$_c", keyB)
+ assert.NotEqual(t, keyA, keyB, "Length-prefixed encoding must guarantee
zero collision")
+}
diff --git a/pkg/datasource/sql/util/escape.go
b/pkg/datasource/sql/util/escape.go
index 1970477a..adf7dc4a 100644
--- a/pkg/datasource/sql/util/escape.go
+++ b/pkg/datasource/sql/util/escape.go
@@ -19,6 +19,7 @@ package util
import (
"database/sql"
+ "fmt"
"strings"
"seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
@@ -187,21 +188,25 @@ func DataValidationAndGoOn(sqlUndoLog undo.SQLUndoLog,
conn *sql.Conn) bool {
}
func GetOrderedPkList(image *types.RecordImage, row types.RowImage, dbType
types.DBType) ([]types.ColumnImage, error) {
-
+ if image == nil || image.TableMeta == nil {
+ return nil, fmt.Errorf("invalid record image or table meta is
nil")
+ }
pkColumnNameListByOrder := image.TableMeta.GetPrimaryKeyOnlyName()
- pkColumnNameListNoOrder := make([]types.ColumnImage, 0)
- pkFields := make([]types.ColumnImage, 0)
+ rawPks := row.PrimaryKeys(row.Columns)
+ pkColumnNameListNoOrder := make([]types.ColumnImage, 0, len(rawPks))
+ pkFields := make([]types.ColumnImage, 0, len(pkColumnNameListByOrder))
- for _, column := range row.PrimaryKeys(row.Columns) {
+ for _, column := range rawPks {
column.ColumnName = DelEscape(column.ColumnName, dbType)
pkColumnNameListNoOrder = append(pkColumnNameListNoOrder,
column)
}
for _, pkName := range pkColumnNameListByOrder {
for _, col := range pkColumnNameListNoOrder {
- if strings.Index(col.ColumnName, pkName) > -1 {
+ if strings.EqualFold(col.ColumnName, pkName) {
pkFields = append(pkFields, col)
+ break
}
}
}
diff --git a/pkg/datasource/sql/util/escape_test.go
b/pkg/datasource/sql/util/escape_test.go
index 4571f871..a12ccc71 100644
--- a/pkg/datasource/sql/util/escape_test.go
+++ b/pkg/datasource/sql/util/escape_test.go
@@ -314,3 +314,63 @@ func TestGetOrderedPkListEmptyRow(t *testing.T) {
assert.NotNil(t, result)
assert.Len(t, result, 0)
}
+
+func TestGetOrderedPkListNilGuards(t *testing.T) {
+ row := types.RowImage{
+ Columns: []types.ColumnImage{
+ {ColumnName: "id", Value: 1, KeyType:
types.IndexTypePrimaryKey},
+ },
+ }
+
+ result, err := GetOrderedPkList(nil, row, types.DBTypeMySQL)
+ assert.Error(t, err)
+ assert.Nil(t, result)
+ assert.Contains(t, err.Error(), "invalid record image")
+
+ imageWithNilMeta := &types.RecordImage{
+ TableName: "t_user",
+ TableMeta: nil,
+ }
+ result, err = GetOrderedPkList(imageWithNilMeta, row, types.DBTypeMySQL)
+ assert.Error(t, err)
+ assert.Nil(t, result)
+ assert.Contains(t, err.Error(), "table meta is nil")
+}
+
+func TestGetOrderedPkListCompositePK_ShuffledInput(t *testing.T) {
+ tableMeta := types.TableMeta{
+ TableName: "t_order",
+ ColumnNames: []string{"tenant_id", "order_id", "user_id"},
+ Indexs: map[string]types.IndexMeta{
+ "PRIMARY": {
+ IType: types.IndexTypePrimaryKey,
+ Columns: []types.ColumnMeta{
+ {ColumnName: "tenant_id"},
+ {ColumnName: "order_id"},
+ },
+ },
+ },
+ }
+
+ image := &types.RecordImage{
+ TableName: "t_order",
+ TableMeta: &tableMeta,
+ }
+
+ shuffledRow := types.RowImage{
+ Columns: []types.ColumnImage{
+ {ColumnName: "order_id", Value: 999, KeyType:
types.IndexTypePrimaryKey},
+ {ColumnName: "amount", Value: 88.8, KeyType:
types.IndexTypeNull},
+ {ColumnName: "tenant_id", Value: "tenant_A", KeyType:
types.IndexTypePrimaryKey},
+ },
+ }
+
+ orderedPks, err := GetOrderedPkList(image, shuffledRow,
types.DBTypeMySQL)
+
+ assert.NoError(t, err)
+ assert.Len(t, orderedPks, 2)
+ assert.Equal(t, "tenant_id", orderedPks[0].ColumnName)
+ assert.Equal(t, "tenant_A", orderedPks[0].Value)
+ assert.Equal(t, "order_id", orderedPks[1].ColumnName)
+ assert.Equal(t, 999, orderedPks[1].Value)
+}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]