This is an automated email from the ASF dual-hosted git repository.
terrymanu pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/shardingsphere.git
The following commit(s) were added to refs/heads/master by this push:
new 4004670ef70 Fix incorrect generated key handling for explicit
auto-increment values (#38810)
4004670ef70 is described below
commit 4004670ef7044e7bea76d4d3d29c425d4a8a0b37
Author: somil jain <[email protected]>
AuthorDate: Fri Jul 3 21:20:13 2026 +0530
Fix incorrect generated key handling for explicit auto-increment values
(#38810)
* Fix incorrect generated key handling for explicit auto-increment values
* Updated release notes
* Stop the proxy from asking the backend database for the generated key
* Preserve generated key handling for explicit NULL/0 auto-increment values
* Fix checkstyle error
* Fix spotless error
* Refactor generated key detection using dialect-specific contracts
* Remove unnecessary lines
* Added regression test and handling an edge case
* Handle generated key triggers for INSERT without column list
* Trigger CI
* Trigger CI
* Align generated key option with connector metadata design
* Fix checkstyle error
* Address generated key review feedback
---
RELEASE-NOTES.md | 1 +
...yOption.java => DefaultGeneratedKeyOption.java} | 16 ++-
.../option/keygen/DialectGeneratedKeyOption.java | 22 ++--
.../keygen/DefaultGeneratedKeyOptionTest.java} | 28 +++--
.../metadata/database/MySQLDatabaseMetaData.java | 3 +-
.../database/option/MySQLGeneratedKeyOption.java | 43 +++++++
.../option/MySQLGeneratedKeyOptionTest.java | 47 ++++++++
.../proxy/backend/connector/ProxySQLExecutor.java | 85 +++++++++++++-
.../connector/StandardDatabaseProxyConnector.java | 5 +-
.../backend/connector/ProxySQLExecutorTest.java | 123 ++++++++++++++++++++-
.../StandardDatabaseProxyConnectorTest.java | 52 +++++++++
.../callback/ProxyJDBCExecutorCallbackTest.java | 30 +++++
12 files changed, 425 insertions(+), 30 deletions(-)
diff --git a/RELEASE-NOTES.md b/RELEASE-NOTES.md
index bf53184ef33..c8723984008 100644
--- a/RELEASE-NOTES.md
+++ b/RELEASE-NOTES.md
@@ -29,6 +29,7 @@
1. Pipeline: Fix MySQL zero-value temporal binlog decoding with fractional
precision in migration -
[#38629](https://github.com/apache/shardingsphere/pull/38629)
1. Pipeline: Fix escape MySQL JSON binlog control characters -
[#38800](https://github.com/apache/shardingsphere/pull/38800)
1. Sharding: Support ORDER BY MySQL VARBINARY column by wrapping byte[] values
in a Comparable adapter -
[#38699](https://github.com/apache/shardingsphere/pull/38699)
+1. Proxy: Fix incorrect generated key handling for explicit auto-increment
values - [#38810](https://github.com/apache/shardingsphere/pull/38810)
1. Sharding: Fix generated actual index names exceeding database identifier
length limits while preserving legacy generated index name compatibility -
[#38449](https://github.com/apache/shardingsphere/pull/38449)
1. Sharding: Fix AUTO_INTERVAL sharding failure under JVM default locales that
use comma decimal separators -
[#38806](https://github.com/apache/shardingsphere/pull/38806)
1. DistSQL: Fix case-sensitive storage unit matching in `SHOW RULES USED
STORAGE UNIT` - [#38848](https://github.com/apache/shardingsphere/pull/38848)
diff --git
a/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
b/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DefaultGeneratedKeyOption.java
similarity index 76%
copy from
database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
copy to
database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DefaultGeneratedKeyOption.java
index 4b57de3917d..9b7293ce6c6 100644
---
a/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
+++
b/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DefaultGeneratedKeyOption.java
@@ -17,15 +17,23 @@
package
org.apache.shardingsphere.database.connector.core.metadata.database.metadata.option.keygen;
-import lombok.Getter;
import lombok.RequiredArgsConstructor;
/**
- * Dialect generated key option.
+ * Default generated key option.
*/
@RequiredArgsConstructor
-@Getter
-public final class DialectGeneratedKeyOption {
+public final class DefaultGeneratedKeyOption implements
DialectGeneratedKeyOption {
private final String columnName;
+
+ @Override
+ public String getColumnName() {
+ return columnName;
+ }
+
+ @Override
+ public boolean isGeneratedKeyTriggerValue(final Object value) {
+ return false;
+ }
}
diff --git
a/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
b/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
index 4b57de3917d..ff531deaf58 100644
---
a/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
+++
b/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
@@ -17,15 +17,23 @@
package
org.apache.shardingsphere.database.connector.core.metadata.database.metadata.option.keygen;
-import lombok.Getter;
-import lombok.RequiredArgsConstructor;
-
/**
* Dialect generated key option.
*/
-@RequiredArgsConstructor
-@Getter
-public final class DialectGeneratedKeyOption {
+public interface DialectGeneratedKeyOption {
+
+ /**
+ * Get generated key column name.
+ *
+ * @return generated key column name
+ */
+ String getColumnName();
- private final String columnName;
+ /**
+ * Check if the explicit value triggers an auto-increment key generation.
+ *
+ * @param value explicit insert value
+ * @return whether the value triggers generated key
+ */
+ boolean isGeneratedKeyTriggerValue(Object value);
}
diff --git
a/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
b/database/connector/core/src/test/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DefaultGeneratedKeyOptionTest.java
similarity index 53%
copy from
database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
copy to
database/connector/core/src/test/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DefaultGeneratedKeyOptionTest.java
index 4b57de3917d..2d4f7d0199d 100644
---
a/database/connector/core/src/main/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DialectGeneratedKeyOption.java
+++
b/database/connector/core/src/test/java/org/apache/shardingsphere/database/connector/core/metadata/database/metadata/option/keygen/DefaultGeneratedKeyOptionTest.java
@@ -17,15 +17,25 @@
package
org.apache.shardingsphere.database.connector.core.metadata.database.metadata.option.keygen;
-import lombok.Getter;
-import lombok.RequiredArgsConstructor;
+import org.junit.jupiter.api.Test;
-/**
- * Dialect generated key option.
- */
-@RequiredArgsConstructor
-@Getter
-public final class DialectGeneratedKeyOption {
+import static org.hamcrest.MatcherAssert.assertThat;
+import static org.hamcrest.Matchers.is;
+import static org.junit.jupiter.api.Assertions.assertFalse;
+
+class DefaultGeneratedKeyOptionTest {
+
+ @Test
+ void assertGetColumnName() {
+ DefaultGeneratedKeyOption actual = new
DefaultGeneratedKeyOption("GENERATED_KEY");
+ assertThat(actual.getColumnName(), is("GENERATED_KEY"));
+ }
- private final String columnName;
+ @Test
+ void assertIsGeneratedKeyTriggerValue() {
+ DefaultGeneratedKeyOption actual = new
DefaultGeneratedKeyOption("GENERATED_KEY");
+ assertFalse(actual.isGeneratedKeyTriggerValue("DEFAULT"));
+ assertFalse(actual.isGeneratedKeyTriggerValue(null));
+ assertFalse(actual.isGeneratedKeyTriggerValue(0));
+ }
}
diff --git
a/database/connector/dialect/mysql/src/main/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/MySQLDatabaseMetaData.java
b/database/connector/dialect/mysql/src/main/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/MySQLDatabaseMetaData.java
index 44d8f94b945..623112cbed7 100644
---
a/database/connector/dialect/mysql/src/main/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/MySQLDatabaseMetaData.java
+++
b/database/connector/dialect/mysql/src/main/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/MySQLDatabaseMetaData.java
@@ -31,6 +31,7 @@ import
org.apache.shardingsphere.database.connector.core.metadata.database.metad
import
org.apache.shardingsphere.database.connector.core.metadata.database.metadata.option.version.DialectProtocolVersionOption;
import
org.apache.shardingsphere.database.connector.mysql.metadata.database.option.MySQLDataTypeOption;
import
org.apache.shardingsphere.database.connector.mysql.metadata.database.option.MySQLFunctionOption;
+import
org.apache.shardingsphere.database.connector.mysql.metadata.database.option.MySQLGeneratedKeyOption;
import java.sql.Connection;
import java.util.Arrays;
@@ -84,7 +85,7 @@ public final class MySQLDatabaseMetaData implements
DialectDatabaseMetaData {
@Override
public Optional<DialectGeneratedKeyOption> getGeneratedKeyOption() {
- return Optional.of(new DialectGeneratedKeyOption("GENERATED_KEY"));
+ return Optional.of(new MySQLGeneratedKeyOption());
}
@Override
diff --git
a/database/connector/dialect/mysql/src/main/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/option/MySQLGeneratedKeyOption.java
b/database/connector/dialect/mysql/src/main/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/option/MySQLGeneratedKeyOption.java
new file mode 100644
index 00000000000..cc78d5725de
--- /dev/null
+++
b/database/connector/dialect/mysql/src/main/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/option/MySQLGeneratedKeyOption.java
@@ -0,0 +1,43 @@
+/*
+ * 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
org.apache.shardingsphere.database.connector.mysql.metadata.database.option;
+
+import
org.apache.shardingsphere.database.connector.core.metadata.database.metadata.option.keygen.DialectGeneratedKeyOption;
+
+/**
+ * Generated key option of MySQL.
+ */
+public final class MySQLGeneratedKeyOption implements
DialectGeneratedKeyOption {
+
+ @Override
+ public String getColumnName() {
+ return "GENERATED_KEY";
+ }
+
+ @Override
+ public boolean isGeneratedKeyTriggerValue(final Object value) {
+ if (null == value) {
+ return true;
+ }
+ if (value instanceof Number && 0L == ((Number) value).longValue()) {
+ return true;
+ }
+ String valueStr = value.toString();
+ return "0".equals(valueStr) || "NULL".equalsIgnoreCase(valueStr) ||
"DEFAULT".equalsIgnoreCase(valueStr);
+ }
+}
diff --git
a/database/connector/dialect/mysql/src/test/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/option/MySQLGeneratedKeyOptionTest.java
b/database/connector/dialect/mysql/src/test/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/option/MySQLGeneratedKeyOptionTest.java
new file mode 100644
index 00000000000..ccf9bc84dc9
--- /dev/null
+++
b/database/connector/dialect/mysql/src/test/java/org/apache/shardingsphere/database/connector/mysql/metadata/database/option/MySQLGeneratedKeyOptionTest.java
@@ -0,0 +1,47 @@
+/*
+ * 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
org.apache.shardingsphere.database.connector.mysql.metadata.database.option;
+
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertTrue;
+
+class MySQLGeneratedKeyOptionTest {
+
+ @Test
+ void assertIsGeneratedKeyTriggerValue() {
+ MySQLGeneratedKeyOption generatedKeyOption = new
MySQLGeneratedKeyOption();
+
+ assertTrue(generatedKeyOption.isGeneratedKeyTriggerValue(null));
+ assertTrue(generatedKeyOption.isGeneratedKeyTriggerValue("NULL"));
+ assertTrue(generatedKeyOption.isGeneratedKeyTriggerValue("null"));
+
+ assertTrue(generatedKeyOption.isGeneratedKeyTriggerValue(0));
+ assertTrue(generatedKeyOption.isGeneratedKeyTriggerValue(0L));
+ assertTrue(generatedKeyOption.isGeneratedKeyTriggerValue("0"));
+
+ assertTrue(generatedKeyOption.isGeneratedKeyTriggerValue("DEFAULT"));
+ assertTrue(generatedKeyOption.isGeneratedKeyTriggerValue("default"));
+
+ assertFalse(generatedKeyOption.isGeneratedKeyTriggerValue(-3));
+ assertFalse(generatedKeyOption.isGeneratedKeyTriggerValue("-3"));
+ assertFalse(generatedKeyOption.isGeneratedKeyTriggerValue(123));
+ assertFalse(generatedKeyOption.isGeneratedKeyTriggerValue("test"));
+ }
+}
diff --git
a/proxy/backend/core/src/main/java/org/apache/shardingsphere/proxy/backend/connector/ProxySQLExecutor.java
b/proxy/backend/core/src/main/java/org/apache/shardingsphere/proxy/backend/connector/ProxySQLExecutor.java
index d1dec333dc9..4e9a3b9c4e2 100644
---
a/proxy/backend/core/src/main/java/org/apache/shardingsphere/proxy/backend/connector/ProxySQLExecutor.java
+++
b/proxy/backend/core/src/main/java/org/apache/shardingsphere/proxy/backend/connector/ProxySQLExecutor.java
@@ -20,13 +20,17 @@ package org.apache.shardingsphere.proxy.backend.connector;
import com.google.common.base.Strings;
import lombok.Getter;
import
org.apache.shardingsphere.database.connector.core.metadata.database.metadata.DialectDatabaseMetaData;
+import
org.apache.shardingsphere.database.connector.core.metadata.database.metadata.option.keygen.DialectGeneratedKeyOption;
import
org.apache.shardingsphere.database.connector.core.metadata.database.metadata.option.transaction.DialectTransactionOption;
import
org.apache.shardingsphere.database.connector.core.spi.DatabaseTypedSPILoader;
import org.apache.shardingsphere.database.connector.core.type.DatabaseType;
import
org.apache.shardingsphere.database.connector.core.type.DatabaseTypeRegistry;
import
org.apache.shardingsphere.database.exception.core.exception.transaction.TableModifyInTransactionException;
+import
org.apache.shardingsphere.infra.binder.context.segment.insert.keygen.GeneratedKeyContext;
+import
org.apache.shardingsphere.infra.binder.context.segment.insert.values.InsertValueContext;
import
org.apache.shardingsphere.infra.binder.context.segment.table.TablesContext;
import
org.apache.shardingsphere.infra.binder.context.statement.SQLStatementContext;
+import
org.apache.shardingsphere.infra.binder.context.statement.type.dml.InsertStatementContext;
import org.apache.shardingsphere.infra.config.props.ConfigurationPropertyKey;
import org.apache.shardingsphere.infra.exception.ShardingSpherePreconditions;
import org.apache.shardingsphere.infra.executor.kernel.ExecutorEngine;
@@ -60,6 +64,10 @@ import
org.apache.shardingsphere.proxy.backend.session.ConnectionSession;
import
org.apache.shardingsphere.proxy.backend.session.transaction.TransactionStatus;
import org.apache.shardingsphere.proxy.backend.util.TransactionUtils;
import
org.apache.shardingsphere.sql.parser.statement.core.enums.TransactionIsolationLevel;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.ExpressionSegment;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.complex.CommonExpressionSegment;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.ParameterMarkerExpressionSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.SQLStatement;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.ddl.CloseStatement;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.ddl.CursorStatement;
@@ -67,7 +75,6 @@ import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.ddl.DD
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.ddl.FetchStatement;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.ddl.MoveStatement;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.ddl.TruncateStatement;
-import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.dml.InsertStatement;
import org.apache.shardingsphere.sqlfederation.engine.SQLFederationEngine;
import org.apache.shardingsphere.transaction.api.TransactionType;
import org.apache.shardingsphere.transaction.spi.TransactionHook;
@@ -186,7 +193,7 @@ public final class ProxySQLExecutor {
int maxConnectionsSizePerQuery = ProxyContext.getInstance()
.getContextManager().getMetaDataContexts().getMetaData().getProps().<Integer>getValue(ConfigurationPropertyKey.MAX_CONNECTIONS_SIZE_PER_QUERY);
DialectDatabaseMetaData dialectDatabaseMetaData = new
DatabaseTypeRegistry(executionContext.getSqlStatementContext().getSqlStatement().getDatabaseType()).getDialectDatabaseMetaData();
- boolean isReturnGeneratedKeys =
executionContext.getSqlStatementContext().getSqlStatement() instanceof
InsertStatement && dialectDatabaseMetaData.getGeneratedKeyOption().isPresent();
+ boolean isReturnGeneratedKeys =
isReturnGeneratedKeys(executionContext.getSqlStatementContext(),
dialectDatabaseMetaData, executionContext.getQueryContext().getParameters());
return hasRawExecutionRule(rules)
? rawExecute(executionContext, rules,
maxConnectionsSizePerQuery)
: useDriverToExecute(executionContext, rules,
maxConnectionsSizePerQuery, isReturnGeneratedKeys,
SQLExecutorExceptionHandler.isExceptionThrown());
@@ -257,4 +264,78 @@ public final class ProxySQLExecutor {
.flatMap(optional ->
optional.getSaneQueryResult(executionContext.getSqlStatementContext().getSqlStatement(),
originalException));
return executeResult.map(Collections::singletonList).orElseThrow(() ->
originalException);
}
+
+ /**
+ * Judge whether to return generated keys.
+ *
+ * @param sqlStatementContext SQL statement context
+ * @param dialectDatabaseMetaData dialect database meta data
+ * @param params SQL parameters
+ * @return whether to return generated keys
+ */
+ public static boolean isReturnGeneratedKeys(final SQLStatementContext
sqlStatementContext, final DialectDatabaseMetaData dialectDatabaseMetaData,
final List<Object> params) {
+ if (!(sqlStatementContext instanceof InsertStatementContext) ||
!dialectDatabaseMetaData.getGeneratedKeyOption().isPresent()) {
+ return false;
+ }
+ InsertStatementContext insertStatementContext =
(InsertStatementContext) sqlStatementContext;
+ if (!insertStatementContext.getGeneratedKeyContext().isPresent()) {
+ return false;
+ }
+
+ GeneratedKeyContext generatedKeyContext =
insertStatementContext.getGeneratedKeyContext().get();
+ if (generatedKeyContext.isGenerated()) {
+ return true;
+ }
+
+ DialectGeneratedKeyOption generatedKeyOption =
dialectDatabaseMetaData.getGeneratedKeyOption().get();
+ return isReturnGeneratedKeysForExplicit(insertStatementContext,
params, generatedKeyOption, generatedKeyContext.getColumnName());
+ }
+
+ private static boolean isReturnGeneratedKeysForExplicit(final
InsertStatementContext insertStatementContext, final List<Object> params,
+ final
DialectGeneratedKeyOption generatedKeyOption, final String columnName) {
+ int columnIndex = -1;
+ int index = 0;
+ for (String each : insertStatementContext.getColumnNames()) {
+ if (each.equalsIgnoreCase(columnName)) {
+ columnIndex = index;
+ break;
+ }
+ index++;
+ }
+ if (-1 != columnIndex) {
+ for (InsertValueContext each :
insertStatementContext.getInsertValueContexts()) {
+ if (isReturnGeneratedKeysForExplicit(each, columnIndex,
params, generatedKeyOption)) {
+ return true;
+ }
+ }
+ }
+ return false;
+ }
+
+ private static boolean isReturnGeneratedKeysForExplicit(final
InsertValueContext insertValueContext, final int columnIndex, final
List<Object> params,
+ final
DialectGeneratedKeyOption generatedKeyOption) {
+ List<ExpressionSegment> expressions =
insertValueContext.getValueExpressions();
+ if (null != expressions && columnIndex < expressions.size()) {
+ ExpressionSegment expr = expressions.get(columnIndex);
+ return isReturnGeneratedKeysFromExpression(expr, params,
generatedKeyOption);
+ }
+ return false;
+ }
+
+ private static boolean isReturnGeneratedKeysFromExpression(final
ExpressionSegment expr, final List<Object> params, final
DialectGeneratedKeyOption generatedKeyOption) {
+ if (expr instanceof ParameterMarkerExpressionSegment) {
+ int markerIndex = ((ParameterMarkerExpressionSegment)
expr).getParameterMarkerIndex();
+ if (params.size() > markerIndex) {
+ return
generatedKeyOption.isGeneratedKeyTriggerValue(params.get(markerIndex));
+ }
+ return false;
+ }
+ if (expr instanceof LiteralExpressionSegment) {
+ return
generatedKeyOption.isGeneratedKeyTriggerValue(((LiteralExpressionSegment)
expr).getLiterals());
+ }
+ if (expr instanceof CommonExpressionSegment) {
+ return
generatedKeyOption.isGeneratedKeyTriggerValue(expr.getText());
+ }
+ return false;
+ }
}
diff --git
a/proxy/backend/core/src/main/java/org/apache/shardingsphere/proxy/backend/connector/StandardDatabaseProxyConnector.java
b/proxy/backend/core/src/main/java/org/apache/shardingsphere/proxy/backend/connector/StandardDatabaseProxyConnector.java
index 183643fa41d..4be9613db8d 100644
---
a/proxy/backend/core/src/main/java/org/apache/shardingsphere/proxy/backend/connector/StandardDatabaseProxyConnector.java
+++
b/proxy/backend/core/src/main/java/org/apache/shardingsphere/proxy/backend/connector/StandardDatabaseProxyConnector.java
@@ -78,7 +78,6 @@ import
org.apache.shardingsphere.sql.parser.statement.core.statement.attribute.t
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.ddl.CloseStatement;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.ddl.DDLStatement;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.dml.DMLStatement;
-import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.dml.InsertStatement;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.type.dml.SelectStatement;
import org.apache.shardingsphere.sqlfederation.context.SQLFederationContext;
import org.apache.shardingsphere.transaction.api.TransactionType;
@@ -261,7 +260,7 @@ public final class StandardDatabaseProxyConnector
implements DatabaseProxyConnec
private ResponseHeader doExecuteFederation() throws SQLException {
SQLStatement sqlStatement =
queryContext.getSqlStatementContext().getSqlStatement();
DialectDatabaseMetaData dialectDatabaseMetaData = new
DatabaseTypeRegistry(sqlStatement.getDatabaseType()).getDialectDatabaseMetaData();
- boolean isReturnGeneratedKeys = sqlStatement instanceof
InsertStatement && dialectDatabaseMetaData.getGeneratedKeyOption().isPresent();
+ boolean isReturnGeneratedKeys =
ProxySQLExecutor.isReturnGeneratedKeys(queryContext.getSqlStatementContext(),
dialectDatabaseMetaData, queryContext.getParameters());
DatabaseType protocolType = database.getProtocolType();
ProxyJDBCExecutorCallback callback =
ProxyJDBCExecutorCallbackFactory.newInstance(driverType, protocolType,
database.getResourceMetaData(),
sqlStatement, this, isReturnGeneratedKeys,
SQLExecutorExceptionHandler.isExceptionThrown(), true);
@@ -331,7 +330,9 @@ public final class StandardDatabaseProxyConnector
implements DatabaseProxyConnec
? ((InsertStatementContext)
queryContext.getSqlStatementContext()).getGeneratedKeyContext()
: Optional.empty();
Collection<Comparable<?>> autoIncrementGeneratedValues =
generatedKeyContext.filter(GeneratedKeyContext::isSupportAutoIncrement)
+ .filter(GeneratedKeyContext::isGenerated)
.map(GeneratedKeyContext::getGeneratedValues).orElseGet(Collections::emptyList);
+
UpdateResponseHeader result = new
UpdateResponseHeader(queryContext.getSqlStatementContext().getSqlStatement(),
updateResults, autoIncrementGeneratedValues);
if (isNeedAccumulate()) {
result.mergeUpdateCount();
diff --git
a/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/ProxySQLExecutorTest.java
b/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/ProxySQLExecutorTest.java
index ee4c0da001a..35c4c8782f7 100644
---
a/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/ProxySQLExecutorTest.java
+++
b/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/ProxySQLExecutorTest.java
@@ -24,8 +24,11 @@ import
org.apache.shardingsphere.database.connector.core.metadata.database.metad
import
org.apache.shardingsphere.database.connector.core.spi.DatabaseTypedSPILoader;
import org.apache.shardingsphere.database.connector.core.type.DatabaseType;
import
org.apache.shardingsphere.database.exception.core.exception.transaction.TableModifyInTransactionException;
+import
org.apache.shardingsphere.infra.binder.context.segment.insert.keygen.GeneratedKeyContext;
+import
org.apache.shardingsphere.infra.binder.context.segment.insert.values.InsertValueContext;
import
org.apache.shardingsphere.infra.binder.context.segment.table.TablesContext;
import
org.apache.shardingsphere.infra.binder.context.statement.SQLStatementContext;
+import
org.apache.shardingsphere.infra.binder.context.statement.type.dml.InsertStatementContext;
import org.apache.shardingsphere.infra.config.props.ConfigurationPropertyKey;
import org.apache.shardingsphere.infra.config.rule.RuleConfiguration;
import
org.apache.shardingsphere.infra.executor.kernel.model.ExecutionGroupContext;
@@ -38,6 +41,7 @@ import
org.apache.shardingsphere.infra.executor.sql.execute.engine.raw.callback.
import
org.apache.shardingsphere.infra.executor.sql.execute.result.ExecuteResult;
import
org.apache.shardingsphere.infra.executor.sql.prepare.driver.DriverExecutionPrepareEngine;
import
org.apache.shardingsphere.infra.executor.sql.prepare.driver.jdbc.JDBCDriverType;
+import
org.apache.shardingsphere.infra.executor.sql.prepare.driver.jdbc.StatementOption;
import
org.apache.shardingsphere.infra.executor.sql.prepare.raw.RawExecutionPrepareEngine;
import org.apache.shardingsphere.infra.metadata.ShardingSphereMetaData;
import
org.apache.shardingsphere.infra.metadata.database.ShardingSphereDatabase;
@@ -46,6 +50,7 @@ import
org.apache.shardingsphere.infra.rule.ShardingSphereRule;
import org.apache.shardingsphere.infra.rule.attribute.RuleAttributes;
import
org.apache.shardingsphere.infra.rule.attribute.raw.RawExecutionRuleAttribute;
import
org.apache.shardingsphere.infra.session.connection.transaction.TransactionConnectionContext;
+import org.apache.shardingsphere.infra.session.query.QueryContext;
import org.apache.shardingsphere.infra.spi.type.typed.TypedSPILoader;
import org.apache.shardingsphere.mode.manager.ContextManager;
import
org.apache.shardingsphere.proxy.backend.connector.jdbc.executor.ProxyJDBCExecutor;
@@ -55,6 +60,9 @@ import
org.apache.shardingsphere.proxy.backend.context.BackendExecutorContext;
import org.apache.shardingsphere.proxy.backend.context.ProxyContext;
import org.apache.shardingsphere.proxy.backend.session.ConnectionSession;
import
org.apache.shardingsphere.sql.parser.statement.core.enums.TransactionIsolationLevel;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.complex.CommonExpressionSegment;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.LiteralExpressionSegment;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.ParameterMarkerExpressionSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.SimpleTableSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.generic.table.TableNameSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.statement.SQLStatement;
@@ -242,6 +250,7 @@ class ProxySQLExecutorTest {
Arguments.of("dml-insert-mysql-xa-pass",
createInsertStatement(mysqlDatabaseType), TransactionType.XA, true, true,
false));
}
+ @SuppressWarnings("rawtypes")
@ParameterizedTest(name = "{0}")
@MethodSource("executeScenarios")
void assertExecute(final String name, final boolean hasRawExecutionRule,
final SQLStatement sqlStatement, final boolean inTransaction,
@@ -253,7 +262,7 @@ class ProxySQLExecutorTest {
setExecutorField(proxySQLExecutor, "rawExecutor", rawExecutor);
setExecutorField(proxySQLExecutor, "regularExecutor", regularExecutor);
setExecutorField(proxySQLExecutor, "transactionHooks",
Collections.singletonMap(shardingSphereRule, transactionHook));
- ExecutionContext executionContext =
createExecutionContext(sqlStatement);
+ ExecutionContext executionContext = createExecutionContext(name,
sqlStatement, isReturnGeneratedKeys);
ExecuteResult expectedExecuteResult = mock(ExecuteResult.class);
List<ExecuteResult> expected =
Collections.singletonList(expectedExecuteResult);
if (hasRawExecutionRule) {
@@ -272,7 +281,11 @@ class ProxySQLExecutorTest {
ExecutionGroupContext<JDBCExecutionUnit> jdbcExecutionGroupContext =
mock(ExecutionGroupContext.class);
try (
MockedConstruction<DriverExecutionPrepareEngine> ignored =
mockConstruction(DriverExecutionPrepareEngine.class,
- (mock, context) -> when(mock.prepare(anyString(),
eq(executionContext), anyCollection(),
any(ExecutionGroupReportContext.class))).thenReturn(jdbcExecutionGroupContext)))
{
+ (mock, context) -> {
+ StatementOption statementOption =
(StatementOption) context.arguments().get(4);
+
assertThat(statementOption.isReturnGeneratedKeys(), is(isReturnGeneratedKeys));
+ when(mock.prepare(anyString(),
eq(executionContext), anyCollection(),
any(ExecutionGroupReportContext.class))).thenReturn(jdbcExecutionGroupContext);
+ })) {
when(regularExecutor.execute(any(), eq(jdbcExecutionGroupContext),
eq(isReturnGeneratedKeys), anyBoolean())).thenReturn(expected);
assertThat(proxySQLExecutor.execute(executionContext),
is(expected));
}
@@ -289,6 +302,16 @@ class ProxySQLExecutorTest {
return Stream.of(
Arguments.of("execute-with-raw-rule", true,
createCreateTableStatement(mysqlDatabaseType), true, false, false),
Arguments.of("execute-with-driver-and-generated-keys", false,
createInsertStatement(mysqlDatabaseType), true, true, true),
+
Arguments.of("execute-with-driver-and-explicit-keys-nonspecial", false,
createInsertStatement(mysqlDatabaseType), true, true, false),
+
Arguments.of("execute-with-driver-and-explicit-keys-special-null", false,
createInsertStatement(mysqlDatabaseType), true, true, true),
+ Arguments.of("execute-with-driver-param-marker-nonspecial",
false, createInsertStatement(mysqlDatabaseType), true, true, false),
+ Arguments.of("execute-with-driver-param-marker-special-zero",
false, createInsertStatement(mysqlDatabaseType), true, true, true),
+ Arguments.of("execute-with-driver-param-marker-special-null",
false, createInsertStatement(mysqlDatabaseType), true, true, true),
+
Arguments.of("execute-with-driver-and-explicit-keys-special-default", false,
createInsertStatement(mysqlDatabaseType), true, true, true),
+
Arguments.of("execute-with-driver-and-no-column-list-special-default", false,
createInsertStatement(mysqlDatabaseType), true, true, true),
+
Arguments.of("execute-with-driver-and-no-column-list-special-null", false,
createInsertStatement(mysqlDatabaseType), true, true, true),
+
Arguments.of("execute-with-driver-and-no-column-list-nonspecial", false,
createInsertStatement(mysqlDatabaseType), true, true, false),
+
Arguments.of("execute-with-driver-and-explicit-keys-different-case", false,
createInsertStatement(mysqlDatabaseType), true, true, true),
Arguments.of("execute-with-driver-and-no-transaction", false,
createInsertStatement(postgresqlDatabaseType), false, false, false));
}
@@ -301,7 +324,7 @@ class ProxySQLExecutorTest {
setExecutorField(proxySQLExecutor, "rawExecutor", rawExecutor);
setExecutorField(proxySQLExecutor, "regularExecutor", regularExecutor);
setExecutorField(proxySQLExecutor, "transactionHooks",
Collections.singletonMap(shardingSphereRule, transactionHook));
- ExecutionContext executionContext =
createExecutionContext(sqlStatement);
+ ExecutionContext executionContext = createExecutionContext(name,
sqlStatement, false);
SQLException expectedException = new SQLException("mock prepare
failure");
try (MockedStatic<DatabaseTypedSPILoader> mockedDatabaseTypedSPILoader
= mockStatic(DatabaseTypedSPILoader.class, CALLS_REAL_METHODS)) {
mockedDatabaseTypedSPILoader.when(() ->
DatabaseTypedSPILoader.findService(DialectSaneQueryResultEngine.class,
fixtureDatabaseType)).thenReturn(Optional.of(saneQueryResultEngine));
@@ -385,11 +408,101 @@ class ProxySQLExecutorTest {
return result;
}
- private ExecutionContext createExecutionContext(final SQLStatement
sqlStatement) {
- SQLStatementContext sqlStatementContext =
mock(SQLStatementContext.class);
+ private ExecutionContext createExecutionContext(final String name, final
SQLStatement sqlStatement, final boolean isReturnGeneratedKeys) {
+ SQLStatementContext sqlStatementContext;
+ List<Object> params = Collections.emptyList();
+ if (sqlStatement instanceof InsertStatement) {
+ InsertStatementContext insertStatementContext =
mock(InsertStatementContext.class);
+ if
("execute-with-driver-and-explicit-keys-nonspecial".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getColumnNames()).thenReturn(Collections.singletonList("foo_id"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Collections.singletonList(new
LiteralExpressionSegment(0, 0, -3)));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ } else if
("execute-with-driver-and-explicit-keys-special-null".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getColumnNames()).thenReturn(Collections.singletonList("foo_id"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Collections.singletonList(new
LiteralExpressionSegment(0, 0, null)));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ } else if
("execute-with-driver-param-marker-nonspecial".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getColumnNames()).thenReturn(Collections.singletonList("foo_id"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Collections.singletonList(new
ParameterMarkerExpressionSegment(0, 0, 0)));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ params = Collections.singletonList(-3);
+ } else if
("execute-with-driver-param-marker-special-zero".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getColumnNames()).thenReturn(Collections.singletonList("foo_id"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Collections.singletonList(new
ParameterMarkerExpressionSegment(0, 0, 0)));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ params = Collections.singletonList(0);
+ } else if
("execute-with-driver-param-marker-special-null".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getColumnNames()).thenReturn(Collections.singletonList("foo_id"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Collections.singletonList(new
ParameterMarkerExpressionSegment(0, 0, 0)));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ params = Collections.singletonList(null);
+ } else if
("execute-with-driver-and-explicit-keys-special-default".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getColumnNames()).thenReturn(Collections.singletonList("foo_id"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Collections.singletonList(new
CommonExpressionSegment(0, 0, "DEFAULT")));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ } else if
("execute-with-driver-and-no-column-list-special-default".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getInsertColumnNames()).thenReturn(Collections.emptyList());
+
when(insertStatementContext.getColumnNames()).thenReturn(Arrays.asList("foo_id",
"name"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Arrays.asList(new
CommonExpressionSegment(0, 0, "DEFAULT"),
mock(LiteralExpressionSegment.class)));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ } else if
("execute-with-driver-and-no-column-list-special-null".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getInsertColumnNames()).thenReturn(Collections.emptyList());
+
when(insertStatementContext.getColumnNames()).thenReturn(Arrays.asList("foo_id",
"name"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Arrays.asList(new
LiteralExpressionSegment(0, 0, null), mock(LiteralExpressionSegment.class)));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ } else if
("execute-with-driver-and-no-column-list-nonspecial".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getInsertColumnNames()).thenReturn(Collections.emptyList());
+
when(insertStatementContext.getColumnNames()).thenReturn(Arrays.asList("foo_id",
"name"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Arrays.asList(new
LiteralExpressionSegment(0, 0, -3), mock(LiteralExpressionSegment.class)));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ } else if
("execute-with-driver-and-explicit-keys-different-case".equals(name)) {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("FOO_ID", false);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
when(insertStatementContext.getColumnNames()).thenReturn(Collections.singletonList("foo_id"));
+ InsertValueContext insertValueContext =
mock(InsertValueContext.class);
+
when(insertValueContext.getValueExpressions()).thenReturn(Collections.singletonList(new
CommonExpressionSegment(0, 0, "DEFAULT")));
+
when(insertStatementContext.getInsertValueContexts()).thenReturn(Collections.singletonList(insertValueContext));
+ } else {
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("foo_id", isReturnGeneratedKeys);
+
when(insertStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+ }
+ sqlStatementContext = insertStatementContext;
+ } else {
+ sqlStatementContext = mock(SQLStatementContext.class);
+ }
when(sqlStatementContext.getSqlStatement()).thenReturn(sqlStatement);
ExecutionContext result = mock(ExecutionContext.class);
when(result.getSqlStatementContext()).thenReturn(sqlStatementContext);
+ QueryContext queryContext = mock(QueryContext.class);
+ when(queryContext.getParameters()).thenReturn(params);
+ when(result.getQueryContext()).thenReturn(queryContext);
return result;
}
diff --git
a/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/StandardDatabaseProxyConnectorTest.java
b/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/StandardDatabaseProxyConnectorTest.java
index e181e2a2fe8..9343b4fc1a5 100644
---
a/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/StandardDatabaseProxyConnectorTest.java
+++
b/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/StandardDatabaseProxyConnectorTest.java
@@ -833,6 +833,58 @@ class StandardDatabaseProxyConnectorTest {
verify(proxySQLExecutor).getSqlFederationEngine();
}
+ @Test
+ void
assertExecuteWithExplicitAutoIncrementValueShouldNotReturnGeneratedKey() throws
SQLException {
+ InsertStatementContext sqlStatementContext =
mock(InsertStatementContext.class, RETURNS_DEEP_STUBS);
+ InsertStatement insertStatement =
InsertStatement.builder().databaseType(databaseType).build();
+
when(sqlStatementContext.getSqlStatement()).thenReturn(insertStatement);
+
when(sqlStatementContext.getTablesContext().getDatabaseNames()).thenReturn(Collections.emptyList());
+
when(sqlStatementContext.getTablesContext().getSchemaNames()).thenReturn(Collections.emptyList());
+
when(sqlStatementContext.getTablesContext().getTableNames()).thenReturn(Collections.singleton("t_order"));
+
+ GeneratedKeyContext generatedKeyContext = new
GeneratedKeyContext("order_id", false);
+ generatedKeyContext.setSupportAutoIncrement(true);
+ generatedKeyContext.getGeneratedValues().add(-3L);
+
when(sqlStatementContext.getGeneratedKeyContext()).thenReturn(Optional.of(generatedKeyContext));
+
+ DataNodeRuleAttribute dataNodeRuleAttribute =
mock(DataNodeRuleAttribute.class);
+ when(dataNodeRuleAttribute.isNeedAccumulate(any())).thenReturn(true);
+
+ ShardingSphereDatabase database = mockDatabase();
+
when(database.getRuleMetaData().getAttributes(DataNodeRuleAttribute.class)).thenReturn(Collections.singleton(dataNodeRuleAttribute));
+
+ DatabaseProxyConnector engine =
createDatabaseProxyConnector(JDBCDriverType.STATEMENT,
createQueryContext(sqlStatementContext, database));
+ setField(engine, "proxySQLExecutor", mock(ProxySQLExecutor.class,
RETURNS_DEEP_STUBS));
+
+ ExecutionContext executionContext = mock(ExecutionContext.class,
RETURNS_DEEP_STUBS);
+
when(executionContext.getExecutionUnits()).thenReturn(Collections.singletonList(mock(ExecutionUnit.class)));
+
when(executionContext.getSqlStatementContext()).thenReturn(sqlStatementContext);
+
when(executionContext.getRouteContext().getRouteUnits()).thenReturn(Collections.emptyList());
+
+ AdvancedProxySQLExecutor advancedProxySQLExecutor =
mock(AdvancedProxySQLExecutor.class);
+ when(advancedProxySQLExecutor.execute(any(ExecutionContext.class),
any(ContextManager.class), any(ShardingSphereDatabase.class),
any(DatabaseProxyConnector.class)))
+ .thenReturn(Collections.singletonList(new UpdateResult(1,
0L)));
+
+ try (
+ MockedConstruction<KernelProcessor> mockedKernelProcessor =
mockConstruction(KernelProcessor.class,
+ (mock, context) ->
when(mock.generateExecutionContext(any(QueryContext.class),
any(RuleMetaData.class),
any(ConfigurationProperties.class))).thenReturn(executionContext));
+ MockedConstruction<DatabaseTypeRegistry>
mockedDatabaseTypeRegistry = mockConstruction(DatabaseTypeRegistry.class,
+ (mock, context) ->
when(mock.getDialectDatabaseMetaData()).thenReturn(mock(DialectDatabaseMetaData.class)));
+ MockedConstruction<PushDownMetaDataRefreshEngine>
mockedPushDownMetaDataRefreshEngine =
mockConstruction(PushDownMetaDataRefreshEngine.class,
+ (mock, context) ->
when(mock.isNeedRefresh()).thenReturn(true));
+ MockedStatic<ShardingSphereServiceLoader> serviceLoader =
mockStatic(ShardingSphereServiceLoader.class)) {
+ serviceLoader.when(() ->
ShardingSphereServiceLoader.getServiceInstances(AdvancedProxySQLExecutor.class))
+
.thenReturn(Collections.singleton(advancedProxySQLExecutor));
+
+ UpdateResponseHeader actual = (UpdateResponseHeader)
engine.execute();
+ assertThat(actual.getUpdateCount(), is(1L));
+ assertThat(actual.getLastInsertId(), is(0L));
+ assertThat(mockedKernelProcessor.constructed().size(), is(1));
+ assertThat(mockedDatabaseTypeRegistry.constructed().size(), is(1));
+
assertThat(mockedPushDownMetaDataRefreshEngine.constructed().size(), is(1));
+ }
+ }
+
private SQLStatementContext createSQLStatementContext(final SQLStatement
sqlStatement) {
SQLStatementContext result = mock(SQLStatementContext.class,
RETURNS_DEEP_STUBS);
when(result.getSqlStatement()).thenReturn(sqlStatement);
diff --git
a/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/jdbc/executor/callback/ProxyJDBCExecutorCallbackTest.java
b/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/jdbc/executor/callback/ProxyJDBCExecutorCallbackTest.java
index ba3dce655b1..a9cd325962a 100644
---
a/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/jdbc/executor/callback/ProxyJDBCExecutorCallbackTest.java
+++
b/proxy/backend/core/src/test/java/org/apache/shardingsphere/proxy/backend/connector/jdbc/executor/callback/ProxyJDBCExecutorCallbackTest.java
@@ -72,6 +72,7 @@ import static org.mockito.Mockito.RETURNS_DEEP_STUBS;
import static org.mockito.Mockito.lenient;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.mockStatic;
+import static org.mockito.Mockito.never;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
@@ -232,6 +233,35 @@ class ProxyJDBCExecutorCallbackTest {
}
}
+ @Test
+ void
assertExecuteInsertWithExplicitAutoIncrementValueDoesNotCallGetGeneratedKeys()
throws ReflectiveOperationException, SQLException {
+ setContextManager(mock(ContextManager.class));
+ Statement statement = mock(Statement.class);
+ when(statement.getUpdateCount()).thenReturn(1);
+ ProxyJDBCExecutorCallback callback = mockCallback(false, false, false);
+ callback.executeSQL("insert_sql", statement,
ConnectionMode.MEMORY_STRICTLY, databaseType);
+ verify(statement, never()).getGeneratedKeys();
+ }
+
+ @Test
+ void assertExecuteInsertWithGeneratedKeysCallsGetGeneratedKeys() throws
ReflectiveOperationException, SQLException {
+ setContextManager(mock(ContextManager.class));
+ Statement statement = mock(Statement.class);
+ when(statement.getUpdateCount()).thenReturn(1);
+
+ ResultSet resultSet = mock(ResultSet.class);
+ ResultSetMetaData metaData = mock(ResultSetMetaData.class);
+ when(statement.getGeneratedKeys()).thenReturn(resultSet);
+ when(resultSet.next()).thenReturn(true);
+ when(resultSet.getMetaData()).thenReturn(metaData);
+ when(metaData.getColumnType(1)).thenReturn(Types.INTEGER);
+ when(resultSet.getLong(1)).thenReturn(99L);
+
+ ProxyJDBCExecutorCallback callback = mockCallback(true, false, false);
+ callback.executeSQL("insert_sql", statement,
ConnectionMode.MEMORY_STRICTLY, databaseType);
+ verify(statement).getGeneratedKeys();
+ }
+
private void setContextManager(final ContextManager contextManager) throws
ReflectiveOperationException {
Field contextManagerField =
ProxyContext.class.getDeclaredField("contextManager");
if (null == originalContextManager) {