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) {
*
*
* - ROW dereference by integer index
- *
- ARRAY dereference by integer index
+ *
- ARRAY dereference by integer index, preserving the safe operator's indexing base
*
- MAP dereference by string key
*
*
@@ -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);
+ }
+}