This is an automated email from the ASF dual-hosted git repository.
zhaojinchao95 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 13f88edb2aa Support oracle insert into returning into item statement
parse (#38902)
13f88edb2aa is described below
commit 13f88edb2aa54a68787be0440ef38bc166b9cca4
Author: Zhengqiang Duan <[email protected]>
AuthorDate: Wed Jun 24 18:50:15 2026 +0800
Support oracle insert into returning into item statement parse (#38902)
---
RELEASE-NOTES.md | 1 +
.../database/option/OracleFunctionOption.java | 5 ++-
.../database/option/OracleFunctionOptionTest.java | 1 +
.../insert/values/OnDuplicateUpdateContext.java | 40 +++++++++++++++++++-
.../type/dml/InsertStatementBindingContext.java | 3 +-
.../values/OnDuplicateUpdateContextTest.java | 10 ++---
.../src/main/antlr4/imports/oracle/DMLStatement.g4 | 6 ++-
.../visitor/statement/OracleStatementVisitor.java | 24 ++++++------
.../statement/type/OracleDMLStatementVisitor.java | 44 ++++++++++++++++++++++
.../dml/standard/type/SelectStatementAssert.java | 11 +++++-
.../dml/standard/SelectStatementTestCase.java | 4 ++
.../parser/src/main/resources/case/dml/insert.xml | 25 ++++++++++++
.../parser/src/main/resources/case/dml/update.xml | 2 +-
.../main/resources/sql/supported/dml/insert.xml | 1 +
14 files changed, 153 insertions(+), 24 deletions(-)
diff --git a/RELEASE-NOTES.md b/RELEASE-NOTES.md
index 1bd07db9e6c..de91e5a4e71 100644
--- a/RELEASE-NOTES.md
+++ b/RELEASE-NOTES.md
@@ -50,6 +50,7 @@
1. SQL Parser: Support mysql, doris insert & replace rows statement parse -
[#38585](https://github.com/apache/shardingsphere/pull/38585)
1. SQL Parser: Support Oracle model, pivot, XML and hierarchical query parsing
and binding - [#38689](https://github.com/apache/shardingsphere/pull/38689)
1. SQL Parser: Preserve temporal literal text and raw value for MySQL and
Oracle date-time literal parsing -
[#38886](https://github.com/apache/shardingsphere/pull/38886)
+1. SQL Parser: Support oracle insert into returning into item statement parse
- [#38902](https://github.com/apache/shardingsphere/pull/38902)
1. SQL Binder: Support select order by index bind metadata -
[#38386](https://github.com/apache/shardingsphere/pull/38386)
1. SQL Binder: Support SQL bind when with temp table name is same with
physical table - [#38411](https://github.com/apache/shardingsphere/pull/38411)
1. Metadata: Support parsing query properties from Oracle JDBC URLs -
[#38901](https://github.com/apache/shardingsphere/pull/38901)
diff --git
a/database/connector/dialect/oracle/src/main/java/org/apache/shardingsphere/database/connector/oracle/metadata/database/option/OracleFunctionOption.java
b/database/connector/dialect/oracle/src/main/java/org/apache/shardingsphere/database/connector/oracle/metadata/database/option/OracleFunctionOption.java
index 491480070f8..f245c199d7a 100644
---
a/database/connector/dialect/oracle/src/main/java/org/apache/shardingsphere/database/connector/oracle/metadata/database/option/OracleFunctionOption.java
+++
b/database/connector/dialect/oracle/src/main/java/org/apache/shardingsphere/database/connector/oracle/metadata/database/option/OracleFunctionOption.java
@@ -30,8 +30,9 @@ public final class OracleFunctionOption implements
DialectFunctionOption {
// TODO remove ROWNUM_ and ROW_NUMBER, move DAY to anthor method
@duanzhengqiang
private static final Collection<String> UNPARENTHESIZED_FUNCTION_NAMES =
new CaseInsensitiveSet<>(Arrays.asList(
- "CONNECT_BY_ISCYCLE", "CONNECT_BY_ISLEAF", "CURRENT_DATE",
"CURRENT_TIME", "CURRENT_TIMESTAMP", "CURRENT_USER", "CURRVAL", "DAY",
"DBTIMEZONE", "LEVEL", "LOCALTIME",
- "LOCALTIMESTAMP", "NEXTVAL", "ORA_ROWSCN", "ROWID", "ROWNUM",
"ROWNUM_", "ROW_NUMBER", "SESSIONTIMEZONE", "SESSION_USER", "SYSDATE",
"SYSTIMESTAMP", "UID", "USER"));
+ "CONNECT_BY_ISCYCLE", "CONNECT_BY_ISLEAF", "CURRENT_DATE",
"CURRENT_TIME", "CURRENT_TIMESTAMP", "CURRENT_USER", "CURRVAL", "DAY",
"DBTIMEZONE", "DEFAULT", "LEVEL",
+ "LOCALTIME", "LOCALTIMESTAMP", "NEXTVAL", "ORA_ROWSCN", "ROWID",
"ROWNUM", "ROWNUM_", "ROW_NUMBER", "SESSIONTIMEZONE", "SESSION_USER",
"SYSDATE", "SYSTIMESTAMP", "UID",
+ "USER"));
@Override
public String getIfNullFunctionName() {
diff --git
a/database/connector/dialect/oracle/src/test/java/org/apache/shardingsphere/database/connector/oracle/metadata/database/option/OracleFunctionOptionTest.java
b/database/connector/dialect/oracle/src/test/java/org/apache/shardingsphere/database/connector/oracle/metadata/database/option/OracleFunctionOptionTest.java
index c7daa98fd2f..8a239f2b9e1 100644
---
a/database/connector/dialect/oracle/src/test/java/org/apache/shardingsphere/database/connector/oracle/metadata/database/option/OracleFunctionOptionTest.java
+++
b/database/connector/dialect/oracle/src/test/java/org/apache/shardingsphere/database/connector/oracle/metadata/database/option/OracleFunctionOptionTest.java
@@ -43,6 +43,7 @@ class OracleFunctionOptionTest {
assertTrue(functionOption.getUnparenthesizedFunctionNames().contains("CURRVAL"));
assertTrue(functionOption.getUnparenthesizedFunctionNames().contains("DAY"));
assertTrue(functionOption.getUnparenthesizedFunctionNames().contains("DBTIMEZONE"));
+
assertTrue(functionOption.getUnparenthesizedFunctionNames().contains("DEFAULT"));
assertTrue(functionOption.getUnparenthesizedFunctionNames().contains("LEVEL"));
assertTrue(functionOption.getUnparenthesizedFunctionNames().contains("LOCALTIME"));
assertTrue(functionOption.getUnparenthesizedFunctionNames().contains("LOCALTIMESTAMP"));
diff --git
a/infra/binder/core/src/main/java/org/apache/shardingsphere/infra/binder/context/segment/insert/values/OnDuplicateUpdateContext.java
b/infra/binder/core/src/main/java/org/apache/shardingsphere/infra/binder/context/segment/insert/values/OnDuplicateUpdateContext.java
index 294bbcb2550..afeb8a60871 100644
---
a/infra/binder/core/src/main/java/org/apache/shardingsphere/infra/binder/context/segment/insert/values/OnDuplicateUpdateContext.java
+++
b/infra/binder/core/src/main/java/org/apache/shardingsphere/infra/binder/context/segment/insert/values/OnDuplicateUpdateContext.java
@@ -27,10 +27,12 @@ import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.Func
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.extractor.ExpressionExtractor;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.generic.ParameterMarkerSegment;
import java.util.ArrayList;
import java.util.Collection;
import java.util.Collections;
+import java.util.Comparator;
import java.util.List;
import java.util.stream.Collectors;
@@ -48,10 +50,11 @@ public final class OnDuplicateUpdateContext {
private final List<ColumnSegment> columns;
- public OnDuplicateUpdateContext(final Collection<ColumnAssignmentSegment>
assignments, final List<Object> params, final int parametersOffset) {
+ public OnDuplicateUpdateContext(final Collection<ColumnAssignmentSegment>
assignments, final List<Object> params, final int parametersOffset,
+ final Collection<ParameterMarkerSegment>
parameterMarkers) {
List<ExpressionSegment> expressionSegments =
assignments.stream().map(ColumnAssignmentSegment::getValue).collect(Collectors.toList());
valueExpressions = getValueExpressions(expressionSegments);
- parameterMarkerExpressions =
ExpressionExtractor.getParameterMarkerExpressions(expressionSegments);
+ parameterMarkerExpressions =
getParameterMarkerExpressions(expressionSegments, assignments,
parameterMarkers);
parameterCount = parameterMarkerExpressions.size();
parameters = getParameters(params, parametersOffset);
columns = assignments.stream().map(each ->
each.getColumns().get(0)).collect(Collectors.toList());
@@ -63,6 +66,39 @@ public final class OnDuplicateUpdateContext {
return result;
}
+ private List<ParameterMarkerExpressionSegment>
getParameterMarkerExpressions(final Collection<ExpressionSegment>
expressionSegments,
+
final Collection<ColumnAssignmentSegment> assignments,
+
final Collection<ParameterMarkerSegment> parameterMarkers) {
+ List<ParameterMarkerExpressionSegment> result =
ExpressionExtractor.getParameterMarkerExpressions(expressionSegments);
+ for (ParameterMarkerSegment each : parameterMarkers) {
+ if (isInAssignments(each, assignments) &&
!containsParameterMarker(result, each)) {
+ result.add(each instanceof ParameterMarkerExpressionSegment
+ ? (ParameterMarkerExpressionSegment) each
+ : new
ParameterMarkerExpressionSegment(each.getStartIndex(), each.getStopIndex(),
each.getParameterIndex()));
+ }
+ }
+
result.sort(Comparator.comparingInt(ParameterMarkerExpressionSegment::getParameterMarkerIndex));
+ return result;
+ }
+
+ private boolean isInAssignments(final ParameterMarkerSegment
parameterMarkerSegment, final Collection<ColumnAssignmentSegment> assignments) {
+ for (ColumnAssignmentSegment each : assignments) {
+ if (each.getStartIndex() <= parameterMarkerSegment.getStartIndex()
&& each.getStopIndex() >= parameterMarkerSegment.getStopIndex()) {
+ return true;
+ }
+ }
+ return false;
+ }
+
+ private boolean containsParameterMarker(final
Collection<ParameterMarkerExpressionSegment> parameterMarkerExpressions, final
ParameterMarkerSegment parameterMarkerSegment) {
+ for (ParameterMarkerExpressionSegment each :
parameterMarkerExpressions) {
+ if (each.getParameterMarkerIndex() ==
parameterMarkerSegment.getParameterIndex()) {
+ return true;
+ }
+ }
+ return false;
+ }
+
private List<Object> getParameters(final List<Object> params, final int
paramsOffset) {
if (params.isEmpty() || 0 == parameterCount) {
return Collections.emptyList();
diff --git
a/infra/binder/core/src/main/java/org/apache/shardingsphere/infra/binder/context/statement/type/dml/InsertStatementBindingContext.java
b/infra/binder/core/src/main/java/org/apache/shardingsphere/infra/binder/context/statement/type/dml/InsertStatementBindingContext.java
index 125ee0b69a4..3159fe63029 100644
---
a/infra/binder/core/src/main/java/org/apache/shardingsphere/infra/binder/context/statement/type/dml/InsertStatementBindingContext.java
+++
b/infra/binder/core/src/main/java/org/apache/shardingsphere/infra/binder/context/statement/type/dml/InsertStatementBindingContext.java
@@ -121,7 +121,8 @@ public final class InsertStatementBindingContext implements
SQLStatementContext
return Optional.empty();
}
Collection<ColumnAssignmentSegment> onDuplicateKeyColumns =
onDuplicateKeyColumnsSegment.get().getColumns();
- OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(onDuplicateKeyColumns, params, parametersOffset.get());
+ OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(onDuplicateKeyColumns, params, parametersOffset.get(),
+ baseContext.getSqlStatement().getParameterMarkers());
parametersOffset.addAndGet(onDuplicateUpdateContext.getParameterCount());
return Optional.of(onDuplicateUpdateContext);
}
diff --git
a/infra/binder/core/src/test/java/org/apache/shardingsphere/infra/binder/context/segment/insert/values/OnDuplicateUpdateContextTest.java
b/infra/binder/core/src/test/java/org/apache/shardingsphere/infra/binder/context/segment/insert/values/OnDuplicateUpdateContextTest.java
index 3d48ad3825d..f466c692832 100644
---
a/infra/binder/core/src/test/java/org/apache/shardingsphere/infra/binder/context/segment/insert/values/OnDuplicateUpdateContextTest.java
+++
b/infra/binder/core/src/test/java/org/apache/shardingsphere/infra/binder/context/segment/insert/values/OnDuplicateUpdateContextTest.java
@@ -41,7 +41,7 @@ class OnDuplicateUpdateContextTest {
@Test
void assertInstanceConstructedOk() throws NoSuchMethodException,
InvocationTargetException, IllegalAccessException {
- OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(Collections.emptyList(), Collections.emptyList(), 0);
+ OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(Collections.emptyList(), Collections.emptyList(), 0,
Collections.emptyList());
assertThat(onDuplicateUpdateContext.getValueExpressions(),
is(Plugins.getMemberAccessor()
.invoke(OnDuplicateUpdateContext.class.getDeclaredMethod("getValueExpressions",
Collection.class), onDuplicateUpdateContext, Collections.emptyList())));
assertThat(onDuplicateUpdateContext.getParameters(),
is(Plugins.getMemberAccessor()
@@ -54,7 +54,7 @@ class OnDuplicateUpdateContextTest {
String parameterValue1 = "test1";
String parameterValue2 = "test2";
List<Object> params = Arrays.asList(parameterValue1, parameterValue2);
- OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(assignments, params, 0);
+ OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(assignments, params, 0, Collections.emptyList());
Object valueFromInsertValueContext1 =
onDuplicateUpdateContext.getValue(0);
assertThat(valueFromInsertValueContext1, is(parameterValue1));
Object valueFromInsertValueContext2 =
onDuplicateUpdateContext.getValue(1);
@@ -73,7 +73,7 @@ class OnDuplicateUpdateContextTest {
void assertGetValueWhenLiteralExpressionSegment() {
Object literalObject = new Object();
Collection<ColumnAssignmentSegment> assignments =
createLiteralExpressionSegment(literalObject);
- OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(assignments, Collections.emptyList(), 0);
+ OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(assignments, Collections.emptyList(), 0,
Collections.emptyList());
Object valueFromInsertValueContext =
onDuplicateUpdateContext.getValue(0);
assertThat(valueFromInsertValueContext, is(literalObject));
}
@@ -104,7 +104,7 @@ class OnDuplicateUpdateContextTest {
void assertGetColumn() {
Object literalObject = new Object();
Collection<ColumnAssignmentSegment> assignments =
createLiteralExpressionSegment(literalObject);
- OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(assignments, Collections.emptyList(), 0);
+ OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(assignments, Collections.emptyList(), 0,
Collections.emptyList());
ColumnSegment column = onDuplicateUpdateContext.getColumn(0);
assertThat(column,
is(assignments.iterator().next().getColumns().get(0)));
}
@@ -115,7 +115,7 @@ class OnDuplicateUpdateContextTest {
createAssignmentSegment(createBinaryOperationExpression()),
createAssignmentSegment(new
ParameterMarkerExpressionSegment(0, 10, 5)),
createAssignmentSegment(new LiteralExpressionSegment(0, 10,
new Object())));
- OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(assignments, Arrays.asList("1", "2"), 0);
+ OnDuplicateUpdateContext onDuplicateUpdateContext = new
OnDuplicateUpdateContext(assignments, Arrays.asList("1", "2"), 0,
Collections.emptyList());
assertThat(onDuplicateUpdateContext.getParameterCount(), is(2));
}
}
diff --git
a/parser/sql/engine/dialect/oracle/src/main/antlr4/imports/oracle/DMLStatement.g4
b/parser/sql/engine/dialect/oracle/src/main/antlr4/imports/oracle/DMLStatement.g4
index 8d62abc0a81..8e556cce318 100644
---
a/parser/sql/engine/dialect/oracle/src/main/antlr4/imports/oracle/DMLStatement.g4
+++
b/parser/sql/engine/dialect/oracle/src/main/antlr4/imports/oracle/DMLStatement.g4
@@ -57,7 +57,11 @@ insertValuesClause
;
returningClause
- : (RETURN | RETURNING) exprs INTO dataItem (COMMA_ dataItem)*
+ : (RETURN | RETURNING) exprs INTO returningIntoItem (COMMA_
returningIntoItem)*
+ ;
+
+returningIntoItem
+ : dataItem | parameterMarker
;
dmlTableExprClause
diff --git
a/parser/sql/engine/dialect/oracle/src/main/java/org/apache/shardingsphere/sql/parser/engine/oracle/visitor/statement/OracleStatementVisitor.java
b/parser/sql/engine/dialect/oracle/src/main/java/org/apache/shardingsphere/sql/parser/engine/oracle/visitor/statement/OracleStatementVisitor.java
index 697f2eb1a2c..59ae5044e63 100644
---
a/parser/sql/engine/dialect/oracle/src/main/java/org/apache/shardingsphere/sql/parser/engine/oracle/visitor/statement/OracleStatementVisitor.java
+++
b/parser/sql/engine/dialect/oracle/src/main/java/org/apache/shardingsphere/sql/parser/engine/oracle/visitor/statement/OracleStatementVisitor.java
@@ -782,12 +782,7 @@ public abstract class OracleStatementVisitor extends
OracleStatementBaseVisitor<
return new LiteralExpressionSegment(context.start.getStartIndex(),
context.stop.getStopIndex(), ((BooleanLiteralValue) astNode).getValue());
}
if (astNode instanceof ParameterMarkerValue) {
- ParameterMarkerValue parameterMarker = (ParameterMarkerValue)
astNode;
- ParameterMarkerExpressionSegment segment = new
ParameterMarkerExpressionSegment(context.start.getStartIndex(),
context.stop.getStopIndex(),
- parameterMarker.getValue(), parameterMarker.getType());
- globalParameterMarkerSegments.add(segment);
- statementParameterMarkerSegments.add(segment);
- return segment;
+ return
createParameterMarkerExpressionSegment((ParameterMarkerValue) astNode,
context.start.getStartIndex(), context.stop.getStopIndex());
}
if (astNode instanceof SubquerySegment) {
return new SubqueryExpressionSegment((SubquerySegment) astNode);
@@ -806,11 +801,7 @@ public abstract class OracleStatementVisitor extends
OracleStatementBaseVisitor<
return new SubquerySegment(startIndex, stopIndex,
(SelectStatement) visit(ctx.subquery()), getOriginalText(ctx.subquery()));
}
if (null != ctx.parameterMarker()) {
- ParameterMarkerValue parameterMarker = (ParameterMarkerValue)
visit(ctx.parameterMarker());
- ParameterMarkerExpressionSegment segment = new
ParameterMarkerExpressionSegment(startIndex, stopIndex,
parameterMarker.getValue(), parameterMarker.getType());
- globalParameterMarkerSegments.add(segment);
- statementParameterMarkerSegments.add(segment);
- return segment;
+ return
createParameterMarkerExpressionSegment(ctx.parameterMarker());
}
if (null != ctx.literals()) {
return SQLUtils.createLiteralExpression(visit(ctx.literals()),
startIndex, stopIndex, ctx.literals().start.getInputStream().getText(new
Interval(startIndex, stopIndex)));
@@ -839,6 +830,17 @@ public abstract class OracleStatementVisitor extends
OracleStatementBaseVisitor<
return visitRemainSimpleExpr(ctx, startIndex, stopIndex);
}
+ protected ParameterMarkerExpressionSegment
createParameterMarkerExpressionSegment(final ParameterMarkerContext ctx) {
+ return createParameterMarkerExpressionSegment((ParameterMarkerValue)
visit(ctx), ctx.getStart().getStartIndex(), ctx.getStop().getStopIndex());
+ }
+
+ private ParameterMarkerExpressionSegment
createParameterMarkerExpressionSegment(final ParameterMarkerValue
parameterMarker, final int startIndex, final int stopIndex) {
+ ParameterMarkerExpressionSegment result = new
ParameterMarkerExpressionSegment(startIndex, stopIndex,
parameterMarker.getValue(), parameterMarker.getType());
+ globalParameterMarkerSegments.add(result);
+ statementParameterMarkerSegments.add(result);
+ return result;
+ }
+
private ASTNode visitRemainSimpleExpr(final SimpleExprContext ctx, final
int startIndex, final int stopIndex) {
if (null != ctx.OR_()) {
ExpressionSegment left = (ExpressionSegment)
visit(ctx.simpleExpr(0));
diff --git
a/parser/sql/engine/dialect/oracle/src/main/java/org/apache/shardingsphere/sql/parser/engine/oracle/visitor/statement/type/OracleDMLStatementVisitor.java
b/parser/sql/engine/dialect/oracle/src/main/java/org/apache/shardingsphere/sql/parser/engine/oracle/visitor/statement/type/OracleDMLStatementVisitor.java
index 73a2b993ffc..96a43f9caa2 100644
---
a/parser/sql/engine/dialect/oracle/src/main/java/org/apache/shardingsphere/sql/parser/engine/oracle/visitor/statement/type/OracleDMLStatementVisitor.java
+++
b/parser/sql/engine/dialect/oracle/src/main/java/org/apache/shardingsphere/sql/parser/engine/oracle/visitor/statement/type/OracleDMLStatementVisitor.java
@@ -84,6 +84,8 @@ import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.QueryN
import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.QueryTableExprClauseContext;
import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.QueryTableExprContext;
import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.ReferenceModelContext;
+import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.ReturningClauseContext;
+import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.ReturningIntoItemContext;
import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.RollupCubeClauseContext;
import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.SelectContext;
import
org.apache.shardingsphere.sql.parser.autogen.OracleStatementParser.SelectFromClauseContext;
@@ -117,6 +119,7 @@ import
org.apache.shardingsphere.sql.parser.statement.core.extractor.TableExtrac
import org.apache.shardingsphere.sql.parser.statement.core.enums.CombineType;
import org.apache.shardingsphere.sql.parser.statement.core.enums.JoinType;
import
org.apache.shardingsphere.sql.parser.statement.core.enums.OrderDirection;
+import org.apache.shardingsphere.sql.parser.statement.core.enums.SubqueryType;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dal.VariableSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.assignment.ColumnAssignmentSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.assignment.InsertValuesSegment;
@@ -141,6 +144,7 @@ import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simp
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.simple.SimpleExpressionSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.subquery.SubqueryExpressionSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.expr.subquery.SubquerySegment;
+import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.ReturningSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.item.ColumnProjectionSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.item.DatetimeProjectionSegment;
import
org.apache.shardingsphere.sql.parser.statement.core.segment.dml.item.ExpressionProjectionSegment;
@@ -226,6 +230,9 @@ public final class OracleDMLStatementVisitor extends
OracleStatementVisitor impl
if (null != ctx.whereClause()) {
result.where((WhereSegment) visit(ctx.whereClause()));
}
+ if (null != ctx.returningClause()) {
+ result.returning((ReturningSegment) visit(ctx.returningClause()));
+ }
UpdateStatement updateStatement = result.build();
updateStatement.addParameterMarkers(ctx.getParent() instanceof
ExecuteContext ? getGlobalParameterMarkerSegments() :
popAllStatementParameterMarkerSegments());
updateStatement.getVariableNames().addAll(getVariableNames());
@@ -239,9 +246,11 @@ public final class OracleDMLStatementVisitor extends
OracleStatementVisitor impl
}
if (null != ctx.dmlTableExprClause().dmlSubqueryClause()) {
SubquerySegment subquerySegment = (SubquerySegment)
visit(ctx.dmlTableExprClause().dmlSubqueryClause());
+
subquerySegment.setSelect(subquerySegment.getSelect().withSubqueryType(SubqueryType.TABLE));
return new SubqueryTableSegment(ctx.start.getStartIndex(),
ctx.stop.getStopIndex(), subquerySegment);
}
SubquerySegment subquerySegment = (SubquerySegment)
visit(ctx.dmlTableExprClause().tableCollectionExpr());
+
subquerySegment.setSelect(subquerySegment.getSelect().withSubqueryType(SubqueryType.TABLE));
return new SubqueryTableSegment(ctx.start.getStartIndex(),
ctx.stop.getStopIndex(), subquerySegment);
}
@@ -342,6 +351,7 @@ public final class OracleDMLStatementVisitor extends
OracleStatementVisitor impl
.table(insertStatement.getTable().orElse(null))
.insertColumns(insertStatement.getInsertColumns().orElse(null))
.insertSelect(insertSelect)
+ .returning(null == ctx.returningClause() ? null :
(ReturningSegment) visit(ctx.returningClause()))
.values(insertValues)
.build();
result.getVariableNames().addAll(getVariableNames());
@@ -434,6 +444,9 @@ public final class OracleDMLStatementVisitor extends
OracleStatementVisitor impl
if (null != ctx.whereClause()) {
result.where((WhereSegment) visit(ctx.whereClause()));
}
+ if (null != ctx.returningClause()) {
+ result.returning((ReturningSegment) visit(ctx.returningClause()));
+ }
DeleteStatement deleteStatement = result.build();
deleteStatement.addParameterMarkers(ctx.getParent() instanceof
ExecuteContext ? getGlobalParameterMarkerSegments() :
popAllStatementParameterMarkerSegments());
deleteStatement.getVariableNames().addAll(getVariableNames());
@@ -907,6 +920,37 @@ public final class OracleDMLStatementVisitor extends
OracleStatementVisitor impl
return new BooleanLiteralValue(false);
}
+ @Override
+ public ASTNode visitReturningClause(final ReturningClauseContext ctx) {
+ ProjectionsSegment projections = new
ProjectionsSegment(ctx.exprs().getStart().getStartIndex(),
ctx.exprs().getStop().getStopIndex());
+ for (ExprContext each : ctx.exprs().expr()) {
+ projections.getProjections().add(createReturningProjection(each));
+ }
+ for (ReturningIntoItemContext each : ctx.returningIntoItem()) {
+ if (null != each.parameterMarker()) {
+ createParameterMarkerExpressionSegment(each.parameterMarker());
+ }
+ }
+ return new ReturningSegment(ctx.getStart().getStartIndex(),
ctx.getStop().getStopIndex(), projections);
+ }
+
+ private ProjectionSegment createReturningProjection(final ExprContext ctx)
{
+ ASTNode projection = visit(ctx);
+ if (projection instanceof ProjectionSegment) {
+ return (ProjectionSegment) projection;
+ }
+ if (projection instanceof ComplexExpressionSegment) {
+ return (ProjectionSegment)
createProjectionForComplexExpressionSegment(projection, null);
+ }
+ if (projection instanceof SimpleExpressionSegment) {
+ return createExpressionProjectionSegment((ExpressionSegment)
projection, null);
+ }
+ if (projection instanceof ExpressionSegment) {
+ return (ProjectionSegment)
createProjectionForExpressionSegment(projection, null);
+ }
+ throw new UnsupportedOperationException("Unsupported Returning
Expression");
+ }
+
@Override
public ASTNode visitSelectList(final SelectListContext ctx) {
ProjectionsSegment result = new
ProjectionsSegment(ctx.getStart().getStartIndex(),
ctx.getStop().getStopIndex());
diff --git
a/test/it/parser/src/main/java/org/apache/shardingsphere/test/it/sql/parser/internal/asserts/statement/dml/standard/type/SelectStatementAssert.java
b/test/it/parser/src/main/java/org/apache/shardingsphere/test/it/sql/parser/internal/asserts/statement/dml/standard/type/SelectStatementAssert.java
index 2d9943bcfd5..7012750a27a 100644
---
a/test/it/parser/src/main/java/org/apache/shardingsphere/test/it/sql/parser/internal/asserts/statement/dml/standard/type/SelectStatementAssert.java
+++
b/test/it/parser/src/main/java/org/apache/shardingsphere/test/it/sql/parser/internal/asserts/statement/dml/standard/type/SelectStatementAssert.java
@@ -48,8 +48,8 @@ import
org.apache.shardingsphere.test.it.sql.parser.internal.cases.parser.jaxb.s
import java.util.Optional;
-import static org.hamcrest.Matchers.is;
import static org.hamcrest.MatcherAssert.assertThat;
+import static org.hamcrest.Matchers.is;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
@@ -82,6 +82,7 @@ public final class SelectStatementAssert {
assertModelClause(assertContext, actual, expected);
assertIntoClause(assertContext, actual, expected);
assertOutfileClause(assertContext, actual, expected);
+ assertSubqueryType(assertContext, actual, expected);
}
private static void assertWindowClause(final SQLCaseAssertContext
assertContext, final SelectStatement actual, final SelectStatementTestCase
expected) {
@@ -230,4 +231,12 @@ public final class SelectStatementAssert {
OutfileClauseAssert.assertIs(assertContext, outfileSegment.get(),
expected.getOutfileClause());
}
}
+
+ private static void assertSubqueryType(final SQLCaseAssertContext
assertContext, final SelectStatement actual, final SelectStatementTestCase
expected) {
+ if (null == expected.getSubqueryType()) {
+ return;
+ }
+ assertTrue(actual.getSubqueryType().isPresent(),
assertContext.getText("Actual subquery type should exist."));
+ assertThat(assertContext.getText("Subquery type assertion error: "),
actual.getSubqueryType().get().name(), is(expected.getSubqueryType()));
+ }
}
diff --git
a/test/it/parser/src/main/java/org/apache/shardingsphere/test/it/sql/parser/internal/cases/parser/jaxb/statement/dml/standard/SelectStatementTestCase.java
b/test/it/parser/src/main/java/org/apache/shardingsphere/test/it/sql/parser/internal/cases/parser/jaxb/statement/dml/standard/SelectStatementTestCase.java
index 2e4b0b35816..d35e50eb300 100644
---
a/test/it/parser/src/main/java/org/apache/shardingsphere/test/it/sql/parser/internal/cases/parser/jaxb/statement/dml/standard/SelectStatementTestCase.java
+++
b/test/it/parser/src/main/java/org/apache/shardingsphere/test/it/sql/parser/internal/cases/parser/jaxb/statement/dml/standard/SelectStatementTestCase.java
@@ -34,6 +34,7 @@ import
org.apache.shardingsphere.test.it.sql.parser.internal.cases.parser.jaxb.s
import
org.apache.shardingsphere.test.it.sql.parser.internal.cases.parser.jaxb.segment.impl.window.ExpectedWindowClause;
import
org.apache.shardingsphere.test.it.sql.parser.internal.cases.parser.jaxb.segment.impl.with.ExpectedWithClause;
+import javax.xml.bind.annotation.XmlAttribute;
import javax.xml.bind.annotation.XmlElement;
/**
@@ -46,6 +47,9 @@ public final class SelectStatementTestCase extends
SQLParserTestCase {
@XmlElement
private ExpectedTable from;
+ @XmlAttribute(name = "subquery-type")
+ private String subqueryType;
+
@XmlElement
private final ExpectedProjections projections = new ExpectedProjections();
diff --git a/test/it/parser/src/main/resources/case/dml/insert.xml
b/test/it/parser/src/main/resources/case/dml/insert.xml
index 9230bef95e0..893ad3c940b 100644
--- a/test/it/parser/src/main/resources/case/dml/insert.xml
+++ b/test/it/parser/src/main/resources/case/dml/insert.xml
@@ -5786,6 +5786,31 @@
</values>
</insert>
+ <insert sql-case-id="insert_returning_into_parameter_marker_oracle"
parameters="1, 'foo_user', bnd1">
+ <table name="t_user" start-index="12" stop-index="17" />
+ <columns start-index="19" stop-index="38">
+ <column name="user_id" start-index="20" stop-index="26" />
+ <column name="user_name" start-index="29" stop-index="37" />
+ </columns>
+ <values>
+ <value>
+ <assignment-value>
+ <parameter-marker-expression parameter-index="0"
start-index="48" stop-index="48" />
+ <literal-expression value="1" start-index="48"
stop-index="48" />
+ </assignment-value>
+ <assignment-value>
+ <parameter-marker-expression parameter-index="1"
start-index="51" stop-index="51" />
+ <literal-expression value="foo_user" start-index="51"
stop-index="60" />
+ </assignment-value>
+ </value>
+ </values>
+ <returning start-index="54" stop-index="77" literal-start-index="63"
literal-stop-index="89">
+ <projections start-index="64" stop-index="70"
literal-start-index="73" literal-stop-index="79">
+ <column-projection name="user_id" start-index="64"
stop-index="70" literal-start-index="73" literal-stop-index="79" />
+ </projections>
+ </returning>
+ </insert>
+
<insert sql-case-id="insert_on_conflit_do_update" parameters="1,'init',1">
<table name="t_order" start-index="12" stop-index="18" />
<columns start-index="21" stop-index="32">
diff --git a/test/it/parser/src/main/resources/case/dml/update.xml
b/test/it/parser/src/main/resources/case/dml/update.xml
index 00e56497db5..f0c20e4f43a 100644
--- a/test/it/parser/src/main/resources/case/dml/update.xml
+++ b/test/it/parser/src/main/resources/case/dml/update.xml
@@ -119,7 +119,7 @@
<table start-index="7" stop-index="31">
<subquery-table start-index="7" stop-index="31">
<subquery>
- <select>
+ <select subquery-type="TABLE">
<projections start-index="15" stop-index="15">
<shorthand-projection start-index="15"
stop-index="15" />
</projections>
diff --git a/test/it/parser/src/main/resources/sql/supported/dml/insert.xml
b/test/it/parser/src/main/resources/sql/supported/dml/insert.xml
index f6dd2e0c60a..32e90ae94ae 100644
--- a/test/it/parser/src/main/resources/sql/supported/dml/insert.xml
+++ b/test/it/parser/src/main/resources/sql/supported/dml/insert.xml
@@ -207,5 +207,6 @@
<sql-case id="insert_without_columns_sql92" value="INSERT INTO t_order
VALUES (1, 'ok')" db-types="SQL92" />
<sql-case id="insert_with_default_sql92" value="INSERT INTO
t_order(status) VALUES (DEFAULT)" db-types="SQL92" />
<sql-case id="insert_table_collection_oracle" value="INSERT INTO
TABLE(SELECT t.nested_col FROM t_nested t) VALUES (1)" db-types="Oracle" />
+ <sql-case id="insert_returning_into_parameter_marker_oracle" value="INSERT
INTO t_user (user_id, user_name) VALUES (?, ?) RETURNING user_id INTO ?"
db-types="Oracle" />
<sql-case id="insert_on_conflit_do_update" value="INSERT INTO t_order (
c1, c2, c3 ) VALUES ( ?, now(), ? ) ON CONFLICT ( c1, c3 ) DO UPDATE set c2 =
now() WHERE user_id= ?" db-types="PostgreSQL" />
</sql-cases>