Character and binary values are checked by length because narrowing casts
+ * may truncate them. Exact numeric values use a target cast, which throws on
+ * overflow.
+ *
+ *
You may provide a custom config to convert other nodes that extend
+ * {@link TableModify}.
*
* @see EnumerableRules#ENUMERABLE_TABLE_MODIFICATION_RULE */
public class EnumerableTableModifyRule extends ConverterRule {
@@ -44,6 +93,7 @@ protected EnumerableTableModifyRule(Config config) {
@Override public @Nullable RelNode convert(RelNode rel) {
final TableModify modify = (TableModify) rel;
+ final RelOptCluster cluster = modify.getCluster();
final ModifiableTable modifiableTable =
modify.getTable().unwrap(ModifiableTable.class);
if (modifiableTable == null) {
@@ -51,11 +101,81 @@ protected EnumerableTableModifyRule(Config config) {
}
final RelTraitSet traitSet =
modify.getTraitSet().replace(EnumerableConvention.INSTANCE);
+ RelNode input = convert(modify.getInput(), traitSet);
+ if (modify.isInsert() || modify.isUpdate()) {
+ // INSERT assigns stored columns; UPDATE assigns columns in the SET list.
+ RelDataType assignmentType = modify.isInsert()
+ ? RelOptTableImpl.realRowType(modify.getTable())
+ : modify.getCatalogReader().createTypeFromProjection(
+ modify.getTable().getRowType(),
+ requireNonNull(modify.getUpdateColumnList(), "updateColumnList"));
+ if (modify.isFlattened()) {
+ // TableModify flattens its input, so flatten the target fields too.
+ assignmentType =
+ SqlTypeUtil.flattenRecordType(cluster.getTypeFactory(), assignmentType, null);
+ }
+
+ final RexBuilder rexBuilder = cluster.getRexBuilder();
+ final List projects =
+ new ArrayList<>(rexBuilder.identityProjects(input.getRowType()));
+ final List checks = new ArrayList<>();
+ // UPDATE appends SET values to the old row; INSERT has only new values.
+ final int assignmentOffset = projects.size() - assignmentType.getFieldCount();
+ for (RelDataTypeField field : assignmentType.getFieldList()) {
+ final int sourceOrdinal = assignmentOffset + field.getIndex();
+ final RexNode source = projects.get(sourceOrdinal);
+ final RelDataType targetType = field.getType();
+ final SqlTypeName targetName = targetType.getSqlTypeName();
+ if (SqlTypeUtil.inCharOrBinaryFamilies(targetType)) {
+ if (targetType.getPrecision() < 0) {
+ continue;
+ }
+ // Check character and binary lengths because their casts may truncate.
+ final RexNode length =
+ rexBuilder.makeCall(SqlTypeUtil.inCharFamily(targetType)
+ ? SqlStdOperatorTable.CHAR_LENGTH
+ : SqlStdOperatorTable.OCTET_LENGTH,
+ source);
+ final RexNode fits =
+ rexBuilder.makeCall(SqlStdOperatorTable.LESS_THAN_OR_EQUAL, length,
+ rexBuilder.makeExactLiteral(
+ BigDecimal.valueOf(targetType.getPrecision())));
+ // IS_NOT_FALSE lets NULL pass this length check.
+ final RexNode valid =
+ rexBuilder.makeCall(SqlStdOperatorTable.IS_NOT_FALSE, fits);
+ checks.add(
+ rexBuilder.makeCall(SqlInternalOperators.THROW_UNLESS,
+ valid, rexBuilder.makeLiteral("Value exceeds precision "
+ + targetType.getPrecision() + " of "
+ + targetType.getFullTypeString())));
+ } else {
+ if (targetName == SqlTypeName.DECIMAL) {
+ // A runtime BigDecimal may exceed its declared precision.
+ if (targetType.getPrecision() < 0 || targetType.getScale() < 0) {
+ continue;
+ }
+ } else if (!SqlTypeUtil.isExactNumeric(targetType)
+ || source.getType().getSqlTypeName() == targetName) {
+ continue;
+ }
+ // Exact numeric casts reject overflow and produce the target value.
+ projects.set(sourceOrdinal, rexBuilder.makeCast(targetType, source));
+ }
+ }
+ if (!checks.isEmpty() || !RexUtil.isIdentity(projects, input.getRowType())) {
+ // RexProgram has one condition, so combine all length checks.
+ input =
+ EnumerableCalc.create(
+ input, RexProgram.create(input.getRowType(), projects,
+ RexUtil.composeConjunction(rexBuilder, checks, true),
+ input.getRowType().getFieldNames(), rexBuilder));
+ }
+ }
return new EnumerableTableModify(
- modify.getCluster(), traitSet,
+ cluster, traitSet,
modify.getTable(),
modify.getCatalogReader(),
- convert(modify.getInput(), traitSet),
+ input,
modify.getOperation(),
modify.getUpdateColumnList(),
modify.getSourceExpressionList(),
diff --git a/server/src/test/java/org/apache/calcite/test/ServerTest.java b/server/src/test/java/org/apache/calcite/test/ServerTest.java
index 0be7fb16efa1..ab7bbd9d5271 100644
--- a/server/src/test/java/org/apache/calcite/test/ServerTest.java
+++ b/server/src/test/java/org/apache/calcite/test/ServerTest.java
@@ -16,11 +16,13 @@
*/
package org.apache.calcite.test;
+import org.apache.calcite.avatica.util.ByteString;
import org.apache.calcite.config.CalciteConnectionProperty;
import org.apache.calcite.jdbc.CalciteConnection;
import org.apache.calcite.jdbc.CalcitePrepare;
import org.apache.calcite.schema.Function;
import org.apache.calcite.schema.FunctionParameter;
+import org.apache.calcite.schema.ModifiableTable;
import org.apache.calcite.server.DdlExecutorImpl;
import org.apache.calcite.server.ServerDdlExecutor;
import org.apache.calcite.sql.SqlNode;
@@ -49,6 +51,7 @@
import java.sql.Statement;
import java.sql.Struct;
import java.util.ArrayList;
+import java.util.Collection;
import java.util.List;
import static org.apache.calcite.test.Matchers.isLinux;
@@ -56,6 +59,7 @@
import static org.hamcrest.CoreMatchers.containsString;
import static org.hamcrest.CoreMatchers.is;
import static org.hamcrest.CoreMatchers.notNullValue;
+import static org.hamcrest.CoreMatchers.nullValue;
import static org.hamcrest.MatcherAssert.assertThat;
import static org.junit.jupiter.api.Assertions.assertArrayEquals;
import static org.junit.jupiter.api.Assertions.assertDoesNotThrow;
@@ -64,6 +68,8 @@
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.junit.jupiter.api.Assertions.fail;
+import static java.util.Objects.requireNonNull;
+
/**
* Unit tests for server and DDL.
*/
@@ -82,6 +88,21 @@ static Connection connect() throws SQLException {
.build());
}
+ private static void assertFails(Statement statement, String sql, String message) {
+ final SQLException e =
+ assertThrows(SQLException.class, () -> statement.executeUpdate(sql));
+ assertThat(e.getMessage(), containsString(message));
+ }
+
+ @SuppressWarnings("unchecked")
+ private static Collection