diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/FieldSelectionConverter.java b/isthmus/src/main/java/io/substrait/isthmus/expression/FieldSelectionConverter.java index 4893d1ef5..146eafa34 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/FieldSelectionConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/FieldSelectionConverter.java @@ -12,6 +12,7 @@ import org.apache.calcite.rex.RexLiteral; import org.apache.calcite.rex.RexNode; import org.apache.calcite.sql.SqlKind; +import org.apache.calcite.sql.fun.SqlItemOperator; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -43,7 +44,7 @@ public FieldSelectionConverter(TypeConverter typeConverter) { * * * @@ -65,7 +66,8 @@ public Optional convert( LOGGER .atWarn() .log( - "Found item operator without literal kind/type. This isn't handled well. Reference was {} with toString {}.", + "Found item operator without literal kind/type. This isn't handled well. Reference" + + " was {} with toString {}.", reference.getKind().name(), reference); return Optional.empty(); @@ -90,15 +92,37 @@ public Optional convert( } case ARRAY: { - Optional index = toInt(literal); + if (!(call.getOperator() instanceof SqlItemOperator)) { + return Optional.empty(); + } + SqlItemOperator operator = (SqlItemOperator) call.getOperator(); + // A Substrait list reference returns null for an out-of-range index, so it cannot + // preserve the error behavior of OFFSET or ORDINAL. + if (!operator.safe || (operator.offset != 0 && operator.offset != 1)) { + return Optional.empty(); + } + if (literal instanceof Expression.NullLiteral) { + return nullIfInputCanBeDiscarded(call, input); + } + + Optional index = toLong(literal); if (index.isEmpty()) { return Optional.empty(); } + // Substrait negative offsets count from the end of the list. Calcite treats an index + // below the operator's base as out of range, including zero for one-based ITEM. + if (index.get() < operator.offset) { + return nullIfInputCanBeDiscarded(call, input); + } + long offset = index.get() - operator.offset; + if (offset > Integer.MAX_VALUE) { + return Optional.empty(); + } if (input instanceof FieldReference) { - return Optional.of(((FieldReference) input).dereferenceList(index.get())); + return Optional.of(((FieldReference) input).dereferenceList((int) offset)); } else { - return Optional.of(FieldReference.newListReference(index.get(), input)); + return Optional.of(FieldReference.newListReference((int) offset, input)); } } @@ -121,6 +145,17 @@ public Optional convert( return Optional.empty(); } + private Optional nullIfInputCanBeDiscarded(RexCall call, Expression input) { + // Calcite still evaluates the array operand for a null or out-of-range index. Only literals + // and references into an existing record are safe to omit; even deterministic calls can fail. + if (input instanceof Literal + || (input instanceof FieldReference + && ((FieldReference) input).inputExpression().isEmpty())) { + return Optional.of(ExpressionCreator.typedNull(typeConverter.toSubstrait(call.getType()))); + } + return Optional.empty(); + } + /** * Converts a numeric literal to an integer index. * @@ -141,6 +176,13 @@ private Optional toInt(Expression.Literal l) { return Optional.empty(); } + private Optional toLong(Expression.Literal literal) { + if (literal instanceof Expression.I64Literal) { + return Optional.of(((Expression.I64Literal) literal).value()); + } + return toInt(literal).map(Integer::longValue); + } + /** * Converts a fixed-char literal to a string key. * diff --git a/isthmus/src/test/java/io/substrait/isthmus/NestedStructQueryTest.java b/isthmus/src/test/java/io/substrait/isthmus/NestedStructQueryTest.java index f1400526a..dbcc66592 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/NestedStructQueryTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/NestedStructQueryTest.java @@ -204,7 +204,7 @@ public RelDataType getRowType(RelDataTypeFactory factory) { + " field: 1 # a\n" + " child {\n" + " list_element {\n" - + " offset: 1\n" + + " offset: 0\n" + " }\n" + " }\n" + " }\n" @@ -244,13 +244,13 @@ public RelDataType getRowType(RelDataTypeFactory factory) { + " field: 1 # a\n" + " child {\n" + " list_element {\n" - + " offset: 1\n" + + " offset: 0\n" + " child {\n" + " list_element {\n" - + " offset: 2\n" + + " offset: 1\n" + " child {\n" + " list_element {\n" - + " offset: 3\n" + + " offset: 2\n" + " }\n" + " }\n" + " }\n" @@ -300,7 +300,7 @@ public RelDataType getRowType(RelDataTypeFactory factory) { + " field: 0 # .b\n" + " child {\n" + " list_element {\n" - + " offset: 2\n" + + " offset: 1\n" + " child {\n" + " struct_field {\n" + " field: 0 # .c\n" diff --git a/isthmus/src/test/java/io/substrait/isthmus/expression/FieldSelectionConverterTest.java b/isthmus/src/test/java/io/substrait/isthmus/expression/FieldSelectionConverterTest.java new file mode 100644 index 000000000..bbb89a6a6 --- /dev/null +++ b/isthmus/src/test/java/io/substrait/isthmus/expression/FieldSelectionConverterTest.java @@ -0,0 +1,226 @@ +package io.substrait.isthmus.expression; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import io.substrait.expression.Expression; +import io.substrait.expression.FieldReference; +import io.substrait.expression.proto.ExpressionProtoConverter; +import io.substrait.extension.ExtensionCollector; +import io.substrait.type.TypeCreator; +import io.substrait.util.EmptyVisitationContext; +import java.math.BigDecimal; +import java.util.List; +import java.util.stream.Stream; +import org.apache.calcite.DataContexts; +import org.apache.calcite.jdbc.JavaTypeFactoryImpl; +import org.apache.calcite.rel.type.RelDataType; +import org.apache.calcite.rex.RexBuilder; +import org.apache.calcite.rex.RexExecutable; +import org.apache.calcite.rex.RexExecutorImpl; +import org.apache.calcite.rex.RexNode; +import org.apache.calcite.sql.SqlOperator; +import org.apache.calcite.sql.fun.SqlLibraryOperators; +import org.apache.calcite.sql.fun.SqlStdOperatorTable; +import org.apache.calcite.sql.type.SqlTypeName; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.MethodSource; +import org.junit.jupiter.params.provider.ValueSource; + +class FieldSelectionConverterTest { + private final JavaTypeFactoryImpl typeFactory = new JavaTypeFactoryImpl(); + private final RexBuilder rexBuilder = new RexBuilder(typeFactory); + private final RelDataType intType = typeFactory.createSqlType(SqlTypeName.INTEGER); + private final RexExpressionConverter converter = new RexExpressionConverter(); + + @ParameterizedTest + @CsvSource({"1, 0, 10", "2, 1, 20", "3, 2, 30", "4, 3,", "0,,", "-1,,", "-2,,"}) + void itemPreservesCalciteIndexing(int index, Integer expectedOffset, Integer expectedValue) { + RexNode call = rexBuilder.makeCall(SqlStdOperatorTable.ITEM, array(), integer(index)); + RexExecutable executable = + RexExecutorImpl.getExecutable(rexBuilder, List.of(call), typeFactory.builder().build()); + executable.setDataContext(DataContexts.EMPTY); + assertEquals(expectedValue, executable.execute()[0]); + + Expression converted = call.accept(converter); + if (expectedOffset == null) { + assertEquals( + TypeCreator.NULLABLE.I32, + assertInstanceOf(Expression.NullLiteral.class, converted).getType()); + } else { + assertEquals(expectedOffset.intValue(), listOffset(converted)); + } + } + + @ParameterizedTest + @MethodSource("safeIndexing") + void safeOperatorsUseTheirOwnBase(SqlOperator operator, long index, int expectedOffset) { + Expression expression = + rexBuilder.makeCall(operator, array(), integer(index)).accept(converter); + assertEquals(expectedOffset, listOffset(expression)); + + // The same offset is used when dereferencing an array column rather than an expression. + RexNode column = rexBuilder.makeInputRef(array().getType(), 0); + Expression columnExpression = + rexBuilder.makeCall(operator, column, integer(index)).accept(converter); + io.substrait.proto.Expression proto = toProto(columnExpression); + assertEquals( + expectedOffset, + proto + .getSelection() + .getDirectReference() + .getStructField() + .getChild() + .getListElement() + .getOffset()); + } + + private static Stream safeIndexing() { + return Stream.of( + Arguments.of(SqlStdOperatorTable.ITEM, 1L, 0), + Arguments.of(SqlStdOperatorTable.ITEM, 2147483648L, Integer.MAX_VALUE), + Arguments.of(SqlLibraryOperators.SAFE_OFFSET, 0L, 0), + Arguments.of(SqlLibraryOperators.SAFE_OFFSET, 1L, 1), + Arguments.of(SqlLibraryOperators.SAFE_OFFSET, 3L, 3), + Arguments.of(SqlLibraryOperators.SAFE_OFFSET, (long) Integer.MAX_VALUE, Integer.MAX_VALUE), + Arguments.of(SqlLibraryOperators.SAFE_ORDINAL, 1L, 0), + Arguments.of(SqlLibraryOperators.SAFE_ORDINAL, 3L, 2), + Arguments.of(SqlLibraryOperators.SAFE_ORDINAL, 2147483648L, Integer.MAX_VALUE)); + } + + @ParameterizedTest + @MethodSource("invalidLowIndexes") + void invalidLowIndexesAreNull(SqlOperator operator, long index) { + Expression converted = rexBuilder.makeCall(operator, array(), integer(index)).accept(converter); + assertEquals( + TypeCreator.NULLABLE.I32, + assertInstanceOf(Expression.NullLiteral.class, converted).getType()); + } + + private static Stream invalidLowIndexes() { + return Stream.of( + Arguments.of(SqlStdOperatorTable.ITEM, 0L), + Arguments.of(SqlStdOperatorTable.ITEM, Long.MIN_VALUE), + Arguments.of(SqlLibraryOperators.SAFE_OFFSET, -1L), + Arguments.of(SqlLibraryOperators.SAFE_OFFSET, Long.MIN_VALUE), + Arguments.of(SqlLibraryOperators.SAFE_ORDINAL, 0L), + Arguments.of(SqlLibraryOperators.SAFE_ORDINAL, -1L)); + } + + @Test + void nullIndexIsNull() { + Expression converted = + rexBuilder + .makeCall(SqlStdOperatorTable.ITEM, array(), rexBuilder.makeNullLiteral(intType)) + .accept(converter); + assertEquals( + TypeCreator.NULLABLE.I32, + assertInstanceOf(Expression.NullLiteral.class, converted).getType()); + } + + @ParameterizedTest + @MethodSource("nullOrInvalidIndexes") + void rejectsDiscardingThrowingArray(SqlOperator operator, Integer index) { + RexNode cast = + rexBuilder.makeAbstractCast(intType, rexBuilder.makeLiteral("not an integer"), false); + RexNode throwingArray = rexBuilder.makeCall(SqlStdOperatorTable.ARRAY_VALUE_CONSTRUCTOR, cast); + RexNode nestedArray = + rexBuilder.makeCall(SqlStdOperatorTable.ARRAY_VALUE_CONSTRUCTOR, throwingArray); + RexNode nestedSelection = + rexBuilder.makeCall(SqlStdOperatorTable.ITEM, nestedArray, integer(1)); + + for (RexNode input : List.of(throwingArray, nestedSelection)) { + RexNode call = rexBuilder.makeCall(operator, input, nullableIndex(index)); + RexExecutable executable = + RexExecutorImpl.getExecutable(rexBuilder, List.of(call), typeFactory.builder().build()); + executable.setDataContext(DataContexts.EMPTY); + assertThrows(NumberFormatException.class, executable::execute); + assertThrows(IllegalArgumentException.class, () -> call.accept(converter)); + } + } + + @ParameterizedTest + @MethodSource("nullOrInvalidIndexes") + void rejectsDiscardingNonliteralArray(SqlOperator operator, Integer index) { + RexNode text = rexBuilder.makeInputRef(typeFactory.createSqlType(SqlTypeName.VARCHAR, 20), 0); + RexNode cast = rexBuilder.makeAbstractCast(intType, text, false); + RexNode array = rexBuilder.makeCall(SqlStdOperatorTable.ARRAY_VALUE_CONSTRUCTOR, cast); + RexNode call = rexBuilder.makeCall(operator, array, nullableIndex(index)); + + assertThrows(IllegalArgumentException.class, () -> call.accept(converter)); + } + + @ParameterizedTest + @MethodSource("nullOrInvalidIndexes") + void literalAndColumnArraysCanBeDiscarded(SqlOperator operator, Integer index) { + RexNode column = rexBuilder.makeInputRef(array().getType(), 0); + for (RexNode input : List.of(array(), column)) { + Expression converted = + rexBuilder.makeCall(operator, input, nullableIndex(index)).accept(converter); + assertEquals( + TypeCreator.NULLABLE.I32, + assertInstanceOf(Expression.NullLiteral.class, converted).getType()); + } + } + + private static Stream nullOrInvalidIndexes() { + return Stream.of( + Arguments.of(SqlStdOperatorTable.ITEM, 0), + Arguments.of(SqlStdOperatorTable.ITEM, -1), + Arguments.of(SqlStdOperatorTable.ITEM, null), + Arguments.of(SqlLibraryOperators.SAFE_OFFSET, -1), + Arguments.of(SqlLibraryOperators.SAFE_OFFSET, null), + Arguments.of(SqlLibraryOperators.SAFE_ORDINAL, 0), + Arguments.of(SqlLibraryOperators.SAFE_ORDINAL, null)); + } + + private RexNode nullableIndex(Integer index) { + return index == null ? rexBuilder.makeNullLiteral(intType) : integer(index); + } + + @ParameterizedTest + @ValueSource(longs = {2147483649L, 4294967297L, Long.MAX_VALUE}) + void unrepresentableOffsetsAreRejected(long index) { + RexNode call = rexBuilder.makeCall(SqlStdOperatorTable.ITEM, array(), integer(index)); + assertThrows(IllegalArgumentException.class, () -> call.accept(converter)); + } + + @ParameterizedTest + @MethodSource("unsafeOperators") + void throwingOperatorsAreRejected(SqlOperator operator) { + RexNode call = rexBuilder.makeCall(operator, array(), integer(4)); + assertThrows(IllegalArgumentException.class, () -> call.accept(converter)); + } + + private static Stream unsafeOperators() { + return Stream.of(SqlLibraryOperators.OFFSET, SqlLibraryOperators.ORDINAL); + } + + private RexNode array() { + return rexBuilder.makeCall( + SqlStdOperatorTable.ARRAY_VALUE_CONSTRUCTOR, integer(10), integer(20), integer(30)); + } + + private RexNode integer(long value) { + RelDataType type = + value >= Integer.MIN_VALUE && value <= Integer.MAX_VALUE + ? intType + : typeFactory.createSqlType(SqlTypeName.BIGINT); + return rexBuilder.makeExactLiteral(BigDecimal.valueOf(value), type); + } + + private int listOffset(Expression expression) { + assertInstanceOf(FieldReference.class, expression); + return toProto(expression).getSelection().getDirectReference().getListElement().getOffset(); + } + + private io.substrait.proto.Expression toProto(Expression expression) { + return expression.accept( + new ExpressionProtoConverter(new ExtensionCollector(), null), + EmptyVisitationContext.INSTANCE); + } +}