diff --git a/core/src/main/java/io/substrait/expression/FieldReference.java b/core/src/main/java/io/substrait/expression/FieldReference.java index 92e726fe3..e6a7a21f2 100644 --- a/core/src/main/java/io/substrait/expression/FieldReference.java +++ b/core/src/main/java/io/substrait/expression/FieldReference.java @@ -161,10 +161,10 @@ public FieldReference dereferenceStruct(int index) { private FieldReference dereference(Type newType, ReferenceSegment nextSegment) { return ImmutableFieldReference.builder() + .from(this) .type(newType) - .addSegments(nextSegment) + .segments(Collections.singletonList(nextSegment)) .addAllSegments(segments()) - .inputExpression(inputExpression()) .build(); } diff --git a/core/src/main/java/io/substrait/expression/proto/ProtoExpressionConverter.java b/core/src/main/java/io/substrait/expression/proto/ProtoExpressionConverter.java index f1e11c0b0..91fd114ad 100644 --- a/core/src/main/java/io/substrait/expression/proto/ProtoExpressionConverter.java +++ b/core/src/main/java/io/substrait/expression/proto/ProtoExpressionConverter.java @@ -89,6 +89,10 @@ public FieldReference from(io.substrait.proto.Expression.FieldReference referenc rootType, getDirectReferenceSegments(reference.getDirectReference())); case OUTER_REFERENCE: { + if (reference.getDirectReference().getStructField().hasChild()) { + throw new UnsupportedOperationException( + "Nested field access in outer references is not yet supported"); + } io.substrait.proto.Expression.FieldReference.OuterReference outerReference = reference.getOuterReference(); int field = reference.getDirectReference().getStructField().getField(); diff --git a/core/src/test/java/io/substrait/expression/FieldReferenceDereferenceTest.java b/core/src/test/java/io/substrait/expression/FieldReferenceDereferenceTest.java new file mode 100644 index 000000000..4038f23bb --- /dev/null +++ b/core/src/test/java/io/substrait/expression/FieldReferenceDereferenceTest.java @@ -0,0 +1,87 @@ +package io.substrait.expression; + +import static org.junit.jupiter.api.Assertions.assertEquals; + +import io.substrait.TestBase; +import io.substrait.type.Type; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.EnumSource; + +class FieldReferenceDereferenceTest extends TestBase { + + enum ReferenceScope { + ROOT, + EXPRESSION, + OUTER_STEPS, + OUTER_ANCHOR, + LAMBDA_CURRENT, + LAMBDA_OUTER + } + + @ParameterizedTest + @EnumSource(ReferenceScope.class) + void structDereferencePreservesScope(ReferenceScope scope) { + FieldReference reference = reference(scope, R.struct(R.BOOLEAN, N.I64)); + + assertDereference( + reference, reference.dereferenceStruct(0), R.BOOLEAN, FieldReference.StructField.of(0)); + } + + @ParameterizedTest + @EnumSource(ReferenceScope.class) + void listDereferencePreservesScope(ReferenceScope scope) { + FieldReference reference = reference(scope, R.list(N.I64)); + + assertDereference( + reference, reference.dereferenceList(2), N.I64, FieldReference.ListElement.of(2)); + } + + @ParameterizedTest + @EnumSource(ReferenceScope.class) + void mapDereferencePreservesScope(ReferenceScope scope) { + FieldReference reference = reference(scope, R.map(R.STRING, N.I64)); + Expression.Literal key = ExpressionCreator.string(false, "key"); + + assertDereference( + reference, reference.dereferenceMap(key), N.I64, FieldReference.MapKey.of(key)); + } + + private FieldReference reference(ReferenceScope scope, Type type) { + return switch (scope) { + case ROOT -> FieldReference.newRootStructReference(1, type); + case EXPRESSION -> + FieldReference.newStructReference( + 1, + Expression.DynamicParameter.builder() + .type(R.struct(R.BOOLEAN, type)) + .parameterReference(0) + .build()); + case OUTER_STEPS -> FieldReference.newRootStructOuterReference(1, type, 2); + case OUTER_ANCHOR -> FieldReference.newRootStructOuterReferenceByRelReference(1, type, 7); + case LAMBDA_CURRENT -> FieldReference.newLambdaParameterReference(0, 1, type); + case LAMBDA_OUTER -> FieldReference.newLambdaParameterReference(2, 1, type); + }; + } + + private void assertDereference( + FieldReference original, + FieldReference dereferenced, + Type expectedType, + FieldReference.ReferenceSegment nextSegment) { + assertEquals( + ImmutableFieldReference.copyOf(original) + .withType(expectedType) + .withSegments(nextSegment, original.segments().get(0)), + dereferenced); + + io.substrait.proto.Expression.FieldReference originalProto = + expressionProtoConverter.toProto(original).getSelection(); + io.substrait.proto.Expression.FieldReference dereferencedProto = + expressionProtoConverter.toProto(dereferenced).getSelection(); + assertEquals(originalProto.getRootTypeCase(), dereferencedProto.getRootTypeCase()); + assertEquals(originalProto.getOuterReference(), dereferencedProto.getOuterReference()); + assertEquals( + originalProto.getLambdaParameterReference(), + dereferencedProto.getLambdaParameterReference()); + } +} diff --git a/core/src/test/java/io/substrait/type/proto/OuterReferenceRoundtripTest.java b/core/src/test/java/io/substrait/type/proto/OuterReferenceRoundtripTest.java index 672467fc7..5c58068b3 100644 --- a/core/src/test/java/io/substrait/type/proto/OuterReferenceRoundtripTest.java +++ b/core/src/test/java/io/substrait/type/proto/OuterReferenceRoundtripTest.java @@ -12,6 +12,8 @@ import io.substrait.expression.proto.ProtoExpressionConverter; import io.substrait.type.Type; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; /** * Round-trip tests for the two outer-reference resolution mechanisms introduced with Substrait @@ -55,6 +57,25 @@ void idBasedOuterReference() { verifyOuterReferenceRoundTrip(reference); } + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void nestedOuterReferenceIsRejected(boolean byAnchor) { + Type.Struct nestedType = R.struct(R.BOOLEAN, R.I64); + FieldReference outer = + byAnchor + ? FieldReference.newRootStructOuterReferenceByRelReference(1, nestedType, 42) + : FieldReference.newRootStructOuterReference(1, nestedType, 1); + io.substrait.proto.Expression proto = + expressionProtoConverter.toProto(outer.dereferenceStruct(0)); + + assertEquals( + "Nested field access in outer references is not yet supported", + assertThrows( + UnsupportedOperationException.class, + () -> protoExpressionConverterWithRoot.from(proto)) + .getMessage()); + } + /** * The two outer-reference forms map to a single protobuf {@code oneof} and are therefore mutually * exclusive: a {@link FieldReference} may not carry both at once. diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java b/isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java index 55e22bc63..850340660 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java @@ -831,6 +831,10 @@ public RexNode visit(FieldReference expr, Context context) throws RuntimeExcepti return rexInputRef; } else if (expr.isOuterReference()) { + if (expr.segments().size() > 1) { + throw new UnsupportedOperationException( + "Nested field access in outer references is not yet supported"); + } final ReferenceSegment segment = expr.segments().get(0); if (segment instanceof FieldReference.StructField) { @@ -852,6 +856,10 @@ public RexNode visit(FieldReference expr, Context context) throws RuntimeExcepti throw new IllegalArgumentException("Unhandled type: " + segment); } } else if (expr.isLambdaParameterReference()) { + if (expr.segments().size() > 1) { + throw new UnsupportedOperationException( + "Nested field access in lambda parameters is not yet supported"); + } // as of now calcite doesn't support nested lambda functions // https://github.com/substrait-io/substrait-java/issues/711 int stepsOut = expr.lambdaParameterReferenceStepsOut().get(); diff --git a/isthmus/src/test/java/io/substrait/isthmus/LambdaExpressionTest.java b/isthmus/src/test/java/io/substrait/isthmus/LambdaExpressionTest.java index fb33407b8..476f8a8df 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/LambdaExpressionTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/LambdaExpressionTest.java @@ -1,5 +1,6 @@ package io.substrait.isthmus; +import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertThrows; import io.substrait.expression.Expression; @@ -37,6 +38,19 @@ void validFieldIndex() { assertFullRoundTrip(project); } + @Test + void nestedLambdaParameterFieldIsRejected() { + Expression.Lambda lambda = + lb.lambda( + List.of(R.struct(R.I32, R.I32), R.I32), params -> params.ref(0).dereferenceStruct(1)); + Project project = Project.builder().addExpressions(lambda).input(emptyTable).build(); + + assertEquals( + "Nested field access in lambda parameters is not yet supported", + assertThrows(UnsupportedOperationException.class, () -> substraitToCalcite.convert(project)) + .getMessage()); + } + // (x: i32) -> 42 @Test void lambdaWithLiteralBody() { diff --git a/isthmus/src/test/java/io/substrait/isthmus/SubqueryPlanTest.java b/isthmus/src/test/java/io/substrait/isthmus/SubqueryPlanTest.java index eedbbc105..f651d0b82 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/SubqueryPlanTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/SubqueryPlanTest.java @@ -8,11 +8,14 @@ import com.google.protobuf.util.JsonFormat; import io.substrait.isthmus.expression.RexExpressionConverter; +import io.substrait.isthmus.sql.SubstraitCreateStatementParser; +import io.substrait.plan.ProtoPlanConverter; import io.substrait.proto.Expression; import io.substrait.proto.Expression.Subquery.SetPredicate.PredicateOp; import io.substrait.proto.FilterRel; import io.substrait.proto.Plan; import java.io.IOException; +import org.apache.calcite.prepare.CalciteCatalogReader; import org.apache.calcite.rel.core.CorrelationId; import org.apache.calcite.rex.RexBuilder; import org.apache.calcite.rex.RexFieldAccess; @@ -26,6 +29,61 @@ class SubqueryPlanTest extends PlanTestBase { // TODO: Add a roundtrip test once the ProtoRelConverter is committed and updated to support // subqueries + @Test + void nestedOuterFieldKeepsItsCorrelationAnchor() throws SqlParseException { + CalciteCatalogReader catalog = + SubstraitCreateStatementParser.processCreateStatementsToCatalog( + "CREATE TABLE outer_table (id INTEGER NOT NULL, s ROW(v INTEGER NOT NULL) NOT NULL);" + + "CREATE TABLE inner_table (id INTEGER NOT NULL, s ROW(v INTEGER NOT NULL) NOT NULL)"); + io.substrait.plan.Plan pojo = + toSubstraitPlan( + "SELECT o.id FROM outer_table o WHERE EXISTS" + + " (SELECT 1 FROM inner_table i WHERE i.id = o.s.v)", + catalog); + Plan plan = toProto(pojo); + + FilterRel outerFilter = + plan.getRelations(0).getRoot().getInput().getProject().getInput().getFilter(); + FilterRel innerFilter = + outerFilter.getCondition().getSubquery().getSetPredicate().getTuples().getFilter(); + Expression.FieldReference outerField = + innerFilter.getCondition().getScalarFunction().getArguments(1).getValue().getSelection(); + + assertTrue(outerFilter.getInput().getRead().getCommon().hasRelAnchor()); + assertTrue(outerField.hasOuterReference()); + assertTrue(outerField.getOuterReference().hasRelReference()); + assertEquals( + outerFilter.getInput().getRead().getCommon().getRelAnchor(), + outerField.getOuterReference().getRelReference()); + assertEquals(1, outerField.getDirectReference().getStructField().getField()); + assertTrue(outerField.getDirectReference().getStructField().hasChild()); + assertEquals( + 0, outerField.getDirectReference().getStructField().getChild().getStructField().getField()); + assertTrue( + innerFilter + .getCondition() + .getScalarFunction() + .getArguments(0) + .getValue() + .getSelection() + .hasRootReference()); + + assertEquals( + "Nested field access in outer references is not yet supported", + assertThrows( + UnsupportedOperationException.class, + () -> new ProtoPlanConverter(extensions).from(plan)) + .getMessage()); + assertEquals( + "Nested field access in outer references is not yet supported", + assertThrows( + UnsupportedOperationException.class, + () -> + new SubstraitToCalcite(converterProvider, catalog) + .convert(pojo.getRoots().get(0))) + .getMessage()); + } + @Test void existsCorrelatedSubquery() throws SqlParseException { SqlToSubstrait s = new SqlToSubstrait();