This is an automated email from the ASF dual-hosted git repository.
funky-eyes pushed a commit to branch 2.x
in repository https://gitbox.apache.org/repos/asf/incubator-seata.git
The following commit(s) were added to refs/heads/2.x by this push:
new 386bbb9bbe bugfix: fix PK extraction for batch inserts with SQL
functions (#8124)
386bbb9bbe is described below
commit 386bbb9bbe72e4145ef0f8d6baeba9d7d5665ba5
Author: somil jain <[email protected]>
AuthorDate: Fri Jun 5 16:27:27 2026 +0530
bugfix: fix PK extraction for batch inserts with SQL functions (#8124)
---
changes/en-us/2.x.md | 1 +
.../rm/datasource/exec/BaseInsertExecutor.java | 49 ++++++----
.../datasource/exec/MySQLInsertExecutorTest.java | 107 +++++++++++++++++++++
.../seata/sqlparser/struct/SqlMethodExpr.java | 36 ++++++-
.../druid/mysql/MySQLInsertRecognizer.java | 13 ++-
.../druid/mysql/MySQLInsertRecognizerTest.java | 29 ++++++
6 files changed, 212 insertions(+), 23 deletions(-)
diff --git a/changes/en-us/2.x.md b/changes/en-us/2.x.md
index 2722d5c6c2..fae3ef11bb 100644
--- a/changes/en-us/2.x.md
+++ b/changes/en-us/2.x.md
@@ -50,6 +50,7 @@ Add changes here for all PR submitted to the 2.x branch.
- [[#8106](https://github.com/apache/incubator-seata/pull/8106)] fix NPE
during AOT proxy creation
- [[#8113](https://github.com/apache/incubator-seata/pull/8113)] fix console
export JSON consistency and download issues
- [[#8118](https://github.com/apache/incubator-seata/pull/8118)] Use explicit
columns in rollback validation query
+- [[#8124](https://github.com/apache/incubator-seata/pull/8124)] Fix PK
extraction for batch inserts with SQL functions
### optimize:
diff --git
a/rm-datasource/src/main/java/org/apache/seata/rm/datasource/exec/BaseInsertExecutor.java
b/rm-datasource/src/main/java/org/apache/seata/rm/datasource/exec/BaseInsertExecutor.java
index 557bafc1c3..0b67c48d5f 100644
---
a/rm-datasource/src/main/java/org/apache/seata/rm/datasource/exec/BaseInsertExecutor.java
+++
b/rm-datasource/src/main/java/org/apache/seata/rm/datasource/exec/BaseInsertExecutor.java
@@ -25,6 +25,7 @@ import org.apache.seata.rm.datasource.PreparedStatementProxy;
import org.apache.seata.rm.datasource.StatementProxy;
import org.apache.seata.rm.datasource.sql.struct.TableRecords;
import org.apache.seata.sqlparser.SQLInsertRecognizer;
+import org.apache.seata.sqlparser.SQLParsingException;
import org.apache.seata.sqlparser.SQLRecognizer;
import org.apache.seata.sqlparser.struct.ColumnMeta;
import org.apache.seata.sqlparser.struct.Null;
@@ -150,20 +151,13 @@ public abstract class BaseInsertExecutor<T, S extends
Statement> extends Abstrac
if (insertRows != null && !insertRows.isEmpty()) {
Map<Integer, ArrayList<Object>> parameters =
preparedStatementProxy.getParameters();
final int rowSize = insertRows.size();
- int totalPlaceholderNum = -1;
+ int totalPlaceholderNum = 0;
for (List<Object> row : insertRows) {
// oracle insert sql statement specify
RETURN_GENERATED_KEYS will append :rowid on sql end
// insert parameter count will than the actual +1
if (row.isEmpty()) {
continue;
}
- int currentRowPlaceholderNum = -1;
- for (Object r : row) {
- if (PLACEHOLDER.equals(r)) {
- totalPlaceholderNum += 1;
- currentRowPlaceholderNum += 1;
- }
- }
String pkKey;
int pkIndex;
List<Object> pkValues;
@@ -176,18 +170,15 @@ public abstract class BaseInsertExecutor<T, S extends
Statement> extends Abstrac
pkIndex = entry.getValue();
Object pkValue = row.get(pkIndex);
if (PLACEHOLDER.equals(pkValue)) {
- int currentRowNotPlaceholderNumBeforePkIndex = 0;
- for (int n = 0, len = row.size(); n < len; n++) {
- Object r = row.get(n);
- if (n < pkIndex && !PLACEHOLDER.equals(r)) {
- currentRowNotPlaceholderNumBeforePkIndex++;
- }
- }
- int idx = totalPlaceholderNum
- - currentRowPlaceholderNum
- + pkIndex
- - currentRowNotPlaceholderNumBeforePkIndex;
+ int idx = getIdx(row, pkIndex,
totalPlaceholderNum);
ArrayList<Object> parameter = parameters.get(idx +
1);
+ if (parameter == null) {
+ throw new SQLParsingException(String.format(
+ "Failed to find PreparedStatement
parameter mapping for primary key. "
+ + "Calculated JDBC index: %d.
Total mapped parameters: %d. "
+ + "Please verify your SQL
placeholders match your query parameters.",
+ (idx + 1), parameters.size()));
+ }
pkValues.addAll(parameter);
} else {
pkValues.add(pkValue);
@@ -196,6 +187,13 @@ public abstract class BaseInsertExecutor<T, S extends
Statement> extends Abstrac
pkValuesMap.put(ColumnUtils.delEscape(pkKey,
getDbType()), pkValues);
}
}
+ for (Object r : row) {
+ if (PLACEHOLDER.equals(r)) {
+ totalPlaceholderNum++;
+ } else if (r instanceof SqlMethodExpr) {
+ totalPlaceholderNum += ((SqlMethodExpr)
r).getPlaceholderCount();
+ }
+ }
}
}
} else {
@@ -223,6 +221,19 @@ public abstract class BaseInsertExecutor<T, S extends
Statement> extends Abstrac
return pkValuesMap;
}
+ private static int getIdx(List<Object> row, int pkIndex, int
totalPlaceholderNum) {
+ int placeholdersBeforePkInRow = 0;
+ for (int n = 0; n < pkIndex; n++) {
+ Object r = row.get(n);
+ if (PLACEHOLDER.equals(r)) {
+ placeholdersBeforePkInRow++;
+ } else if (r instanceof SqlMethodExpr) {
+ placeholdersBeforePkInRow += ((SqlMethodExpr)
r).getPlaceholderCount();
+ }
+ }
+ return totalPlaceholderNum + placeholdersBeforePkInRow;
+ }
+
/**
* default get generated keys.
* @return the generate keys
diff --git
a/rm-datasource/src/test/java/org/apache/seata/rm/datasource/exec/MySQLInsertExecutorTest.java
b/rm-datasource/src/test/java/org/apache/seata/rm/datasource/exec/MySQLInsertExecutorTest.java
index dc8e717115..20b4c8c87c 100644
---
a/rm-datasource/src/test/java/org/apache/seata/rm/datasource/exec/MySQLInsertExecutorTest.java
+++
b/rm-datasource/src/test/java/org/apache/seata/rm/datasource/exec/MySQLInsertExecutorTest.java
@@ -21,6 +21,7 @@ import com.alibaba.druid.mock.MockStatementBase;
import com.alibaba.druid.pool.DruidDataSource;
import com.alibaba.druid.sql.SQLUtils;
import com.alibaba.druid.sql.ast.SQLStatement;
+import com.alibaba.druid.sql.dialect.mysql.parser.MySqlStatementParser;
import com.google.common.collect.Lists;
import org.apache.seata.common.exception.ShouldNeverHappenException;
import org.apache.seata.rm.datasource.ConnectionProxy;
@@ -34,6 +35,7 @@ import org.apache.seata.rm.datasource.mock.MockDriver;
import org.apache.seata.rm.datasource.mock.MockResultSet;
import org.apache.seata.rm.datasource.sql.struct.TableRecords;
import org.apache.seata.sqlparser.SQLInsertRecognizer;
+import org.apache.seata.sqlparser.SQLParsingException;
import org.apache.seata.sqlparser.druid.mysql.MySQLInsertRecognizer;
import org.apache.seata.sqlparser.struct.ColumnMeta;
import org.apache.seata.sqlparser.struct.Null;
@@ -62,9 +64,14 @@ import java.util.HashMap;
import java.util.List;
import java.util.Map;
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
+import static org.junit.jupiter.api.Assertions.assertThrows;
+import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.Mockito.anyString;
import static org.mockito.Mockito.doReturn;
import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.spy;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
@@ -821,6 +828,106 @@ public class MySQLInsertExecutorTest {
Assertions.assertEquals(map.get("ID").size(), 3);
}
+ @Test
+ @SuppressWarnings("unchecked")
+ public void testParsePkValuesWithSpatialFunctions() throws Exception {
+ String sql =
+ "INSERT INTO test_table (id, name, geo) VALUES (?, ?,
ST_GeomFromText(CONCAT('POINT(', ?, ' ', ?, ')'))), (?, ?,
ST_GeomFromText(CONCAT('POINT(', ?, ' ', ?, ')')))";
+
+ Map<Integer, ArrayList<Object>> mockParameters = new HashMap<>();
+ mockParameters.put(1, Lists.newArrayList(10));
+ mockParameters.put(2, Lists.newArrayList("A"));
+ mockParameters.put(3, Lists.newArrayList(111));
+ mockParameters.put(4, Lists.newArrayList(222));
+ mockParameters.put(5, Lists.newArrayList(20));
+ mockParameters.put(6, Lists.newArrayList("B"));
+ mockParameters.put(7, Lists.newArrayList(333));
+ mockParameters.put(8, Lists.newArrayList(444));
+
+ PreparedStatementProxy statementProxy =
mock(PreparedStatementProxy.class);
+ when(statementProxy.getParameters()).thenReturn(mockParameters);
+
+ SQLStatement statement = new
MySqlStatementParser(sql).parseStatement();
+ MySQLInsertRecognizer recognizer = new MySQLInsertRecognizer(sql,
statement);
+
+ MySQLInsertExecutor executor = spy(new
MySQLInsertExecutor(statementProxy, (st, args) -> null, recognizer));
+
+ Map<String, Integer> testPkIndexMap = new HashMap<>();
+ testPkIndexMap.put("id", 0);
+ doReturn(testPkIndexMap).when(executor).getPkIndex();
+ doReturn("mysql").when(executor).getDbType();
+
+ Map<String, List<Object>> pkValuesMap =
executor.parsePkValuesFromStatement();
+
+ assertNotNull(pkValuesMap);
+
+ List<Object> pkValues = pkValuesMap.get("id");
+ assertNotNull(pkValues);
+ assertEquals(2, pkValues.size());
+ assertEquals(10, pkValues.get(0));
+ assertEquals(20, pkValues.get(1));
+ }
+
+ @Test
+ public void testParsePkValues_MissingParameter_ThrowsException() {
+ String sql = "INSERT INTO test_table (id, name) VALUES (?, ?)";
+
+ Map<Integer, ArrayList<Object>> mockParameters = new HashMap<>();
+
+ PreparedStatementProxy statementProxy =
mock(PreparedStatementProxy.class);
+ when(statementProxy.getParameters()).thenReturn(mockParameters);
+
+ SQLStatement statement = new
MySqlStatementParser(sql).parseStatement();
+ MySQLInsertRecognizer recognizer = new MySQLInsertRecognizer(sql,
statement);
+
+ MySQLInsertExecutor executor = spy(new
MySQLInsertExecutor(statementProxy, (st, args) -> null, recognizer));
+
+ Map<String, Integer> testPkIndexMap = new HashMap<>();
+ testPkIndexMap.put("id", 0);
+ doReturn(testPkIndexMap).when(executor).getPkIndex();
+ doReturn("mysql").when(executor).getDbType();
+
+ SQLParsingException exception =
assertThrows(SQLParsingException.class, executor::parsePkValuesFromStatement);
+
+ assertTrue(exception.getMessage().contains("Failed to find
PreparedStatement parameter mapping"));
+ }
+
+ @Test
+ @SuppressWarnings("unchecked")
+ public void testParsePkValues_WithSpatialFunctionBeforePk() {
+ String sql =
+ "INSERT INTO test_table (name, geo, id) VALUES (?,
ST_GeomFromText(CONCAT('POINT(', ?, ' ', ?, ')')), ?)";
+
+ Map<Integer, ArrayList<Object>> mockParameters = new HashMap<>();
+ mockParameters.put(1, Lists.newArrayList("A"));
+ mockParameters.put(2, Lists.newArrayList(111));
+ mockParameters.put(3, Lists.newArrayList(222));
+ mockParameters.put(4, Lists.newArrayList(10));
+
+ PreparedStatementProxy statementProxy =
mock(PreparedStatementProxy.class);
+ when(statementProxy.getParameters()).thenReturn(mockParameters);
+
+ SQLStatement statement = new
MySqlStatementParser(sql).parseStatement();
+ MySQLInsertRecognizer recognizer = new MySQLInsertRecognizer(sql,
statement);
+
+ MySQLInsertExecutor executor = spy(new
MySQLInsertExecutor(statementProxy, (st, args) -> null, recognizer));
+
+ Map<String, Integer> testPkIndexMap = new HashMap<>();
+ testPkIndexMap.put("id", 2);
+ doReturn(testPkIndexMap).when(executor).getPkIndex();
+ doReturn("mysql").when(executor).getDbType();
+
+ Map<String, List<Object>> pkValuesMap =
executor.parsePkValuesFromStatement();
+
+ assertNotNull(pkValuesMap);
+
+ List<Object> pkValues = pkValuesMap.get("id");
+ assertNotNull(pkValues);
+ assertEquals(1, pkValues.size());
+
+ assertEquals(10, pkValues.get(0));
+ }
+
private List<String> mockInsertColumns() {
List<String> columns = new ArrayList<>();
columns.add(ID_COLUMN);
diff --git
a/sqlparser/seata-sqlparser-core/src/main/java/org/apache/seata/sqlparser/struct/SqlMethodExpr.java
b/sqlparser/seata-sqlparser-core/src/main/java/org/apache/seata/sqlparser/struct/SqlMethodExpr.java
index 8793ee67d5..3ea1c423c0 100644
---
a/sqlparser/seata-sqlparser-core/src/main/java/org/apache/seata/sqlparser/struct/SqlMethodExpr.java
+++
b/sqlparser/seata-sqlparser-core/src/main/java/org/apache/seata/sqlparser/struct/SqlMethodExpr.java
@@ -22,7 +22,9 @@ package org.apache.seata.sqlparser.struct;
*/
public class SqlMethodExpr {
- private static SqlMethodExpr instance = new SqlMethodExpr();
+ private static final SqlMethodExpr INSTANCE = new SqlMethodExpr(0);
+
+ private final int placeholderCount;
/**
* Get SqlMethodExpr.
@@ -30,13 +32,41 @@ public class SqlMethodExpr {
* @return the SqlMethodExpr
*/
public static SqlMethodExpr get() {
- return instance;
+ return INSTANCE;
}
- private SqlMethodExpr() {}
+ /**
+ * Instantiates a new SqlMethodExpr with a specific placeholder count.
+ *
+ * @param placeholderCount the number of placeholders inside the method
+ */
+ public SqlMethodExpr(int placeholderCount) {
+ this.placeholderCount = placeholderCount;
+ }
+
+ public int getPlaceholderCount() {
+ return placeholderCount;
+ }
@Override
public String toString() {
return "SQL_METHOD";
}
+
+ @Override
+ public boolean equals(Object obj) {
+ if (this == obj) {
+ return true;
+ }
+ if (!(obj instanceof SqlMethodExpr)) {
+ return false;
+ }
+ SqlMethodExpr other = (SqlMethodExpr) obj;
+ return placeholderCount == other.placeholderCount;
+ }
+
+ @Override
+ public int hashCode() {
+ return Integer.hashCode(placeholderCount);
+ }
}
diff --git
a/sqlparser/seata-sqlparser-druid/src/main/java/org/apache/seata/sqlparser/druid/mysql/MySQLInsertRecognizer.java
b/sqlparser/seata-sqlparser-druid/src/main/java/org/apache/seata/sqlparser/druid/mysql/MySQLInsertRecognizer.java
index a95a786c4d..b0733e6fec 100644
---
a/sqlparser/seata-sqlparser-druid/src/main/java/org/apache/seata/sqlparser/druid/mysql/MySQLInsertRecognizer.java
+++
b/sqlparser/seata-sqlparser-druid/src/main/java/org/apache/seata/sqlparser/druid/mysql/MySQLInsertRecognizer.java
@@ -28,6 +28,7 @@ import com.alibaba.druid.sql.ast.statement.SQLExprTableSource;
import com.alibaba.druid.sql.ast.statement.SQLInsertStatement;
import com.alibaba.druid.sql.dialect.mysql.ast.statement.MySqlInsertStatement;
import com.alibaba.druid.sql.dialect.mysql.visitor.MySqlOutputVisitor;
+import com.alibaba.druid.sql.visitor.SQLASTVisitorAdapter;
import org.apache.seata.common.util.CollectionUtils;
import org.apache.seata.sqlparser.SQLInsertRecognizer;
import org.apache.seata.sqlparser.SQLType;
@@ -126,7 +127,17 @@ public class MySQLInsertRecognizer extends
BaseMySQLRecognizer implements SQLIns
} else if (expr instanceof SQLVariantRefExpr) {
row.add(((SQLVariantRefExpr) expr).getName());
} else if (expr instanceof SQLMethodInvokeExpr) {
- row.add(SqlMethodExpr.get());
+ final int[] placeholderCount = {0};
+ expr.accept(new SQLASTVisitorAdapter() {
+ @Override
+ public boolean visit(SQLVariantRefExpr x) {
+ if ("?".equals(x.getName())) {
+ placeholderCount[0]++;
+ }
+ return true;
+ }
+ });
+ row.add(new SqlMethodExpr(placeholderCount[0]));
} else {
if (primaryKeyIndex.contains(i)) {
wrapSQLParsingException(expr);
diff --git
a/sqlparser/seata-sqlparser-druid/src/test/java/org/apache/seata/sqlparser/druid/mysql/MySQLInsertRecognizerTest.java
b/sqlparser/seata-sqlparser-druid/src/test/java/org/apache/seata/sqlparser/druid/mysql/MySQLInsertRecognizerTest.java
index a3214f0e70..7cc76e6c93 100644
---
a/sqlparser/seata-sqlparser-druid/src/test/java/org/apache/seata/sqlparser/druid/mysql/MySQLInsertRecognizerTest.java
+++
b/sqlparser/seata-sqlparser-druid/src/test/java/org/apache/seata/sqlparser/druid/mysql/MySQLInsertRecognizerTest.java
@@ -20,9 +20,11 @@ import com.alibaba.druid.sql.SQLUtils;
import com.alibaba.druid.sql.ast.SQLStatement;
import com.alibaba.druid.sql.ast.statement.SQLInsertStatement;
import com.alibaba.druid.sql.dialect.mysql.ast.expr.MySqlOrderingExpr;
+import com.alibaba.druid.sql.dialect.mysql.parser.MySqlStatementParser;
import org.apache.seata.sqlparser.SQLParsingException;
import org.apache.seata.sqlparser.SQLType;
import org.apache.seata.sqlparser.druid.AbstractRecognizerTest;
+import org.apache.seata.sqlparser.struct.SqlMethodExpr;
import org.apache.seata.sqlparser.util.JdbcConstants;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.Test;
@@ -31,6 +33,10 @@ import java.util.Arrays;
import java.util.Collections;
import java.util.List;
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertInstanceOf;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
+
/**
* The type My sql insert recognizer test.
*
@@ -216,4 +222,27 @@ public class MySQLInsertRecognizerTest extends
AbstractRecognizerTest {
Assertions.assertTrue(insertColumn.contains("`"));
}
}
+
+ @Test
+ public void testGetInsertRowsWithSqlMethodExpr() {
+ String sql =
+ "INSERT INTO test_table (id, name, geo) VALUES (?, ?,
ST_GeomFromText(CONCAT('POINT(', ?, ' ', ?, ')')))";
+ SQLStatement statement = new
MySqlStatementParser(sql).parseStatement();
+ MySQLInsertRecognizer recognizer = new MySQLInsertRecognizer(sql,
statement);
+
+ List<List<Object>> rows =
recognizer.getInsertRows(Collections.emptyList());
+
+ assertNotNull(rows);
+ assertEquals(1, rows.size());
+
+ List<Object> firstRow = rows.get(0);
+ assertEquals(3, firstRow.size());
+
+ assertEquals("?", firstRow.get(0));
+ assertEquals("?", firstRow.get(1));
+
+ Object geoColumn = firstRow.get(2);
+ assertInstanceOf(SqlMethodExpr.class, geoColumn);
+ assertEquals(2, ((SqlMethodExpr) geoColumn).getPlaceholderCount());
+ }
}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]