From 0f9253b6d7367766ad811da3dd0fb6361561fce5 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Sat, 3 Oct 2026 17:55:51 -0400 Subject: [PATCH 1/2] fix!: preserve mandatory read and post-join filters Apply mandatory read predicates before projection and emit, and post-join predicates after join output formation. Decode logical, lateral, hash, and merge post-join references against their direct output schemas. BREAKING CHANGE: Calcite and Spark now enforce embedded mandatory filters. Post-join field references use direct join output before emit instead of concatenated inputs. Regenerate plans that reference removed semi/anti columns or use concatenated-input indices for right-oriented joins; outer-join references now preserve nullable output types. --- .../substrait/relation/ProtoRelConverter.java | 39 +- .../type/proto/JoinRoundtripTest.java | 90 ++++- .../isthmus/SubstraitRelNodeConverter.java | 24 +- .../isthmus/EmbeddedPredicateTest.java | 355 ++++++++++++++++++ .../spark/logical/ToLogicalPlan.scala | 21 +- .../spark/MandatoryPredicatesSuite.scala | 188 ++++++++++ 6 files changed, 689 insertions(+), 28 deletions(-) create mode 100644 isthmus/src/test/java/io/substrait/isthmus/EmbeddedPredicateTest.java create mode 100644 spark/src/test/scala/io/substrait/spark/MandatoryPredicatesSuite.scala diff --git a/core/src/main/java/io/substrait/relation/ProtoRelConverter.java b/core/src/main/java/io/substrait/relation/ProtoRelConverter.java index 716df2291..f36b97cb0 100644 --- a/core/src/main/java/io/substrait/relation/ProtoRelConverter.java +++ b/core/src/main/java/io/substrait/relation/ProtoRelConverter.java @@ -1008,15 +1008,20 @@ protected Join newJoin(JoinRel rel) { Type.Struct unionedStruct = Type.Struct.builder().from(leftStruct).from(rightStruct).build(); ProtoExpressionConverter converter = new ProtoExpressionConverter(lookup, extensions, unionedStruct, this); + Join.JoinType joinType = Join.JoinType.fromProto(rel.getType()); ImmutableJoin.Builder builder = Join.builder() .left(left) .right(right) .condition(converter.from(rel.getExpression())) - .joinType(Join.JoinType.fromProto(rel.getType())) - .postJoinFilter( - Optional.ofNullable( - rel.hasPostJoinFilter() ? converter.from(rel.getPostJoinFilter()) : null)); + .joinType(joinType); + + if (rel.hasPostJoinFilter()) { + ProtoExpressionConverter outputConverter = + new ProtoExpressionConverter( + lookup, extensions, Join.deriveRecordType(joinType, left, right), this); + builder.postJoinFilter(outputConverter.from(rel.getPostJoinFilter())); + } if (rel.hasAdvancedExtension()) { builder.extension(protoExtensionConverter.fromProto(rel.getAdvancedExtension())); @@ -1050,14 +1055,17 @@ protected Rel newLateralJoin(LateralJoinRel rel) { Optional.ofNullable( rel.hasExpression() ? converter.from(rel.getExpression()) : null)) .joinType(Join.JoinType.fromProto(rel.getType())) - .postJoinFilter( - Optional.ofNullable( - rel.hasPostJoinFilter() ? converter.from(rel.getPostJoinFilter()) : null)) // A lateral join validates that it carries an anchor at construction time, so the // anchor has to be set here rather than being left to applyRelCommon, which only runs // after build() (it then sees the same value and skips it). .relAnchor(relAnchor); + if (rel.hasPostJoinFilter()) { + ProtoExpressionConverter outputConverter = + new ProtoExpressionConverter(lookup, extensions, builder.build().getRecordType(), this); + builder.postJoinFilter(outputConverter.from(rel.getPostJoinFilter())); + } + if (rel.hasAdvancedExtension()) { builder.extension(protoExtensionConverter.fromProto(rel.getAdvancedExtension())); } @@ -1126,14 +1134,16 @@ protected Rel newHashJoin(HashJoinRel rel) { .right(right) .keys(comparisonJoinKeys(rel.getKeysList(), leftConverter, rightConverter)) .joinType(HashJoin.JoinType.fromProto(rel.getType())) - .postJoinFilter( - Optional.ofNullable( - rel.hasPostJoinFilter() ? unionConverter.from(rel.getPostJoinFilter()) : null)) .residualExpression( Optional.ofNullable( rel.hasResidualExpression() ? unionConverter.from(rel.getResidualExpression()) : null)); + if (rel.hasPostJoinFilter()) { + ProtoExpressionConverter outputConverter = + new ProtoExpressionConverter(lookup, extensions, builder.build().getRecordType(), this); + builder.postJoinFilter(outputConverter.from(rel.getPostJoinFilter())); + } if (rel.hasAdvancedExtension()) { builder.extension(protoExtensionConverter.fromProto(rel.getAdvancedExtension())); } @@ -1165,15 +1175,18 @@ protected Rel newMergeJoin(MergeJoinRel rel) { .right(right) .keys(comparisonJoinKeys(rel.getKeysList(), leftConverter, rightConverter)) .joinType(MergeJoin.JoinType.fromProto(rel.getType())) - .postJoinFilter( - Optional.ofNullable( - rel.hasPostJoinFilter() ? unionConverter.from(rel.getPostJoinFilter()) : null)) .residualExpression( Optional.ofNullable( rel.hasResidualExpression() ? unionConverter.from(rel.getResidualExpression()) : null)); + if (rel.hasPostJoinFilter()) { + ProtoExpressionConverter outputConverter = + new ProtoExpressionConverter(lookup, extensions, builder.build().getRecordType(), this); + builder.postJoinFilter(outputConverter.from(rel.getPostJoinFilter())); + } + if (rel.hasAdvancedExtension()) { builder.extension(protoExtensionConverter.fromProto(rel.getAdvancedExtension())); } diff --git a/core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java b/core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java index c8c2ddad3..d63ca42db 100644 --- a/core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java +++ b/core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java @@ -10,6 +10,8 @@ import java.util.Arrays; import java.util.List; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.EnumSource; class JoinRoundtripTest extends TestBase { @@ -25,6 +27,39 @@ class JoinRoundtripTest extends TestBase { Arrays.asList("d", "e", "f"), Arrays.asList(R.FP64, R.STRING, R.I64)); + @ParameterizedTest + @EnumSource( + value = Join.JoinType.class, + names = { + "INNER", + "LEFT", + "RIGHT", + "OUTER", + "LEFT_SEMI", + "LEFT_ANTI", + "RIGHT_SEMI", + "RIGHT_ANTI" + }) + void postJoinFilterUsesOutputSchemaBeforeEmit(Join.JoinType joinType) { + Join join = + sb.join( + input -> sb.equal(sb.fieldReference(input, 0), sb.fieldReference(input, 5)), + joinType, + leftTable, + rightTable); + Join filtered = + Join.builder() + .from(join) + .postJoinFilter( + sb.and( + sb.isNull(sb.fieldReference(join, 0)), + sb.isNull(sb.fieldReference(join, join.getRecordType().fields().size() - 1)))) + .remap(sb.remap(0)) + .build(); + + verifyRoundTrip(filtered); + } + @Test void hashJoin() { List leftKeys = Arrays.asList(0, 1); @@ -36,6 +71,21 @@ void hashJoin() { verifyRoundTrip(relWithoutKeys); } + @ParameterizedTest + @EnumSource(value = HashJoin.JoinType.class, names = "UNKNOWN", mode = EnumSource.Mode.EXCLUDE) + void hashPostJoinFilterUsesOutputSchemaBeforeEmit(HashJoin.JoinType joinType) { + HashJoin join = sb.hashJoin(List.of(0), List.of(2), joinType, leftTable, rightTable); + verifyRoundTrip( + HashJoin.builder() + .from(join) + .postJoinFilter( + sb.and( + sb.isNull(sb.fieldReference(join, 0)), + sb.isNull(sb.fieldReference(join, join.getRecordType().fields().size() - 1)))) + .remap(sb.remap(0)) + .build()); + } + @Test void hashJoinWithResidualExpression() { List leftKeys = Arrays.asList(0, 1); @@ -65,6 +115,21 @@ void mergeJoin() { verifyRoundTrip(relWithoutKeys); } + @ParameterizedTest + @EnumSource(value = MergeJoin.JoinType.class, names = "UNKNOWN", mode = EnumSource.Mode.EXCLUDE) + void mergePostJoinFilterUsesOutputSchemaBeforeEmit(MergeJoin.JoinType joinType) { + MergeJoin join = sb.mergeJoin(List.of(0), List.of(2), joinType, leftTable, rightTable); + verifyRoundTrip( + MergeJoin.builder() + .from(join) + .postJoinFilter( + sb.and( + sb.isNull(sb.fieldReference(join, 0)), + sb.isNull(sb.fieldReference(join, join.getRecordType().fields().size() - 1)))) + .remap(sb.remap(0)) + .build()); + } + @Test void mergeJoinWithResidualExpression() { List leftKeys = Arrays.asList(0, 1); @@ -130,10 +195,13 @@ void lateralJoinWithoutCondition() { verifyRoundTrip(rel); } - @Test - void lateralJoinWithAnchorAndPostFilter() { + @ParameterizedTest + @EnumSource( + value = Join.JoinType.class, + names = {"INNER", "LEFT", "LEFT_SEMI", "LEFT_ANTI", "LEFT_SINGLE", "LEFT_MARK"}) + void lateralJoinWithAnchorAndPostFilter(Join.JoinType joinType) { // A lateral join sets a rel anchor so the right input can reference the current left row. - Rel rel = + LateralJoin join = LateralJoin.builder() .left(leftTable) .right(rightTable) @@ -141,13 +209,17 @@ void lateralJoinWithAnchorAndPostFilter() { sb.equal( sb.fieldReference(Arrays.asList(leftTable, rightTable), 0), sb.fieldReference(Arrays.asList(leftTable, rightTable), 5))) - .postJoinFilter( - sb.equal( - sb.fieldReference(Arrays.asList(leftTable, rightTable), 2), - sb.fieldReference(Arrays.asList(leftTable, rightTable), 4))) - .joinType(Join.JoinType.LEFT) + .joinType(joinType) .relAnchor(1) .build(); - verifyRoundTrip(rel); + verifyRoundTrip( + LateralJoin.builder() + .from(join) + .postJoinFilter( + sb.and( + sb.isNull(sb.fieldReference(join, 0)), + sb.isNull(sb.fieldReference(join, join.getRecordType().fields().size() - 1)))) + .remap(sb.remap(0)) + .build()); } } diff --git a/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java b/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java index 24197794c..924cc86b5 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java @@ -205,9 +205,22 @@ public RelNode visit(Filter filter, Context context) throws RuntimeException { @Override public RelNode visit(NamedScan namedScan, Context context) throws RuntimeException { RelNode node = relBuilder.scan(namedScan.getNames()).build(); + node = applyFilter(node, namedScan.getFilter(), context); return applyRelCommon(applyProjection(node, namedScan.getProjection()), namedScan); } + private RelNode applyFilter(RelNode input, Optional filter, Context context) { + if (filter.isEmpty()) { + return input; + } + // Embedded predicates use the operator's direct row, before projection or emit mapping. This + // is an internal input, not another anchored Substrait relation; enclosing scopes still own + // any outer references used by the predicate. + context.enterScope(AnchoredInput.of(Optional.empty(), input.getRowType())); + RexNode condition = filter.get().accept(expressionRexConverter, context); + return relBuilder.push(input).filter(context.exitScope(), condition).build(); + } + @Override public RelNode visit(LocalFiles localFiles, Context context) throws RuntimeException { return visitFallback(localFiles, context); @@ -257,6 +270,7 @@ public RelNode visit(Join join, Context context) throws RuntimeException { JoinRelType joinType = asJoinRelType(join); RelNode node = relBuilder.push(left).push(right).join(joinType, condition, context.exitScope()).build(); + node = applyFilter(node, join.getPostJoinFilter(), context); return applyRelCommon(node, join, left, right); } @@ -1052,7 +1066,10 @@ public RelNode visit(VirtualTableScan virtualTableScan, Context context) { } return applyRelCommon( applyProjection( - LogicalValues.create(relBuilder.getCluster(), rowType, tuplesBuilder.build()), + applyFilter( + LogicalValues.create(relBuilder.getCluster(), rowType, tuplesBuilder.build()), + virtualTableScan.getFilter(), + context), virtualTableScan.getProjection()), virtualTableScan); } else { @@ -1063,7 +1080,10 @@ public RelNode visit(VirtualTableScan virtualTableScan, Context context) { // VirtualTableExpansionRule. return applyRelCommon( applyProjection( - VirtualTable.create(relBuilder.getCluster(), rowType, convertedRows), + applyFilter( + VirtualTable.create(relBuilder.getCluster(), rowType, convertedRows), + virtualTableScan.getFilter(), + context), virtualTableScan.getProjection()), virtualTableScan); } diff --git a/isthmus/src/test/java/io/substrait/isthmus/EmbeddedPredicateTest.java b/isthmus/src/test/java/io/substrait/isthmus/EmbeddedPredicateTest.java new file mode 100644 index 000000000..37251a8af --- /dev/null +++ b/isthmus/src/test/java/io/substrait/isthmus/EmbeddedPredicateTest.java @@ -0,0 +1,355 @@ +package io.substrait.isthmus; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import io.substrait.expression.Expression; +import io.substrait.expression.ExpressionCreator; +import io.substrait.expression.FieldReference; +import io.substrait.expression.MaskExpression; +import io.substrait.extension.ExtensionCollector; +import io.substrait.isthmus.calcite.rel.rules.VirtualTableExpansionRule; +import io.substrait.relation.Join; +import io.substrait.relation.Join.JoinType; +import io.substrait.relation.NamedScan; +import io.substrait.relation.ProtoRelConverter; +import io.substrait.relation.Rel; +import io.substrait.relation.RelProtoConverter; +import io.substrait.relation.VirtualTableScan; +import io.substrait.type.NamedStruct; +import java.util.Arrays; +import java.util.List; +import java.util.stream.Collectors; +import org.apache.calcite.DataContext; +import org.apache.calcite.adapter.java.JavaTypeFactory; +import org.apache.calcite.interpreter.Interpreter; +import org.apache.calcite.jdbc.JavaTypeFactoryImpl; +import org.apache.calcite.linq4j.QueryProvider; +import org.apache.calcite.rel.RelNode; +import org.apache.calcite.rel.core.Filter; +import org.apache.calcite.rel.core.Project; +import org.apache.calcite.rel.core.TableScan; +import org.apache.calcite.rel.core.Values; +import org.apache.calcite.rex.RexInputRef; +import org.apache.calcite.rex.RexSubQuery; +import org.apache.calcite.schema.SchemaPlus; +import org.apache.calcite.tools.Frameworks; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.EnumSource; +import org.junit.jupiter.params.provider.ValueSource; + +class EmbeddedPredicateTest extends PlanTestBase { + + private final MaskExpression projection = + MaskExpression.builder() + .select( + MaskExpression.StructSelect.builder() + .addStructItems(MaskExpression.StructItem.of(2), MaskExpression.StructItem.of(0)) + .build()) + .build(); + + @Test + void namedScanFiltersBeforeProjectionAndEmit() { + NamedScan scan = + sb.namedScan( + List.of("example"), + List.of("id", "keep", "label"), + List.of(R.I32, N.BOOLEAN, R.STRING)); + NamedScan filtered = + NamedScan.builder() + .from(scan) + .filter(sb.fieldReference(scan, 1)) + .projection(projection) + .remap(sb.remap(1)) + .build(); + + RelNode node = substraitToCalcite.convert(filtered); + assertRowMatch(node.getRowType(), R.I32); + while (node instanceof Project) { + node = ((Project) node).getInput(); + } + Filter filter = assertInstanceOf(Filter.class, node); + assertInstanceOf(TableScan.class, filter.getInput()); + assertEquals(1, assertInstanceOf(RexInputRef.class, filter.getCondition()).getIndex()); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void virtualTableFiltersBeforeProjectionAndEmit(boolean computed) { + VirtualTableScan scan = + VirtualTableScan.builder() + .initialSchema( + NamedStruct.of( + List.of("id", "keep", "label"), R.struct(R.I32, R.BOOLEAN, R.STRING))) + .addRows( + ExpressionCreator.nestedStruct( + false, + List.of( + computed ? sb.add(sb.i32(1), sb.i32(1)) : sb.i32(2), + sb.bool(true), + sb.str("kept"))), + ExpressionCreator.nestedStruct( + false, List.of(sb.i32(3), sb.bool(false), sb.str("discarded")))) + .build(); + VirtualTableScan filtered = + VirtualTableScan.builder() + .from(scan) + .filter(sb.fieldReference(scan, 1)) + .projection(projection) + .remap(sb.remap(1)) + .build(); + + assertEquals(List.of(List.of(2)), rows(substraitToCalcite.convert(filtered))); + } + + @Test + void namedScanFiltersBeforeEmit() { + NamedScan scan = + sb.namedScan(List.of("example"), List.of("id", "keep"), List.of(R.I32, N.BOOLEAN)); + NamedScan filtered = + NamedScan.builder() + .from(scan) + .filter(sb.fieldReference(scan, 1)) + .remap(sb.remap(0)) + .build(); + + Project project = assertInstanceOf(Project.class, substraitToCalcite.convert(filtered)); + Filter filter = assertInstanceOf(Filter.class, project.getInput()); + assertInstanceOf(TableScan.class, filter.getInput()); + assertEquals(1, assertInstanceOf(RexInputRef.class, filter.getCondition()).getIndex()); + assertRowMatch(project.getRowType(), R.I32); + } + + @Test + void namedScanFalseAndNullFiltersProduceNoRows() { + NamedScan scan = sb.namedScan(List.of("example"), List.of("id"), List.of(R.I32)); + for (Expression condition : List.of(sb.bool(false), ExpressionCreator.typedNull(N.BOOLEAN))) { + NamedScan filtered = NamedScan.builder().from(scan).filter(condition).build(); + + Values values = assertInstanceOf(Values.class, substraitToCalcite.convert(filtered)); + assertTrue(values.getTuples().isEmpty()); + assertRowMatch(values.getRowType(), R.I32); + } + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void virtualTableFiltersLiteralAndComputedRowsBeforeEmit(boolean computed) { + VirtualTableScan scan = + VirtualTableScan.builder() + .initialSchema(NamedStruct.of(List.of("id", "keep"), R.struct(R.I32, N.BOOLEAN))) + .addRows( + ExpressionCreator.nestedStruct( + false, + List.of( + computed ? sb.add(sb.i32(1), sb.i32(1)) : sb.i32(2), + ExpressionCreator.bool(true, true))), + ExpressionCreator.nestedStruct( + false, List.of(sb.i32(3), ExpressionCreator.bool(true, false))), + ExpressionCreator.nestedStruct( + false, List.of(sb.i32(4), ExpressionCreator.typedNull(N.BOOLEAN)))) + .build(); + VirtualTableScan filtered = + VirtualTableScan.builder() + .from(scan) + .filter(sb.fieldReference(scan, 1)) + .remap(sb.remap(0)) + .build(); + + assertEquals(List.of(List.of(2)), rows(substraitToCalcite.convert(filtered))); + } + + @Test + void bestEffortReadFilterMayBeIgnored() { + VirtualTableScan scan = + VirtualTableScan.builder().from(integers(1, 2)).bestEffortFilter(sb.bool(false)).build(); + + assertEquals(List.of(List.of(1), List.of(2)), rows(substraitToCalcite.convert(scan))); + } + + @ParameterizedTest + @EnumSource( + value = JoinType.class, + names = {"LEFT", "RIGHT", "OUTER"}) + void postJoinFilterSeesNullExtendedRows(JoinType joinType) { + Join join = equalityJoin(joinType); + int nullField = joinType == JoinType.RIGHT ? 0 : 1; + Join filtered = + Join.builder() + .from(join) + .postJoinFilter(sb.isNull(sb.fieldReference(join, nullField))) + .build(); + + RelNode converted = substraitToCalcite.convert(filtered); + Filter filter = assertInstanceOf(Filter.class, converted); + assertInstanceOf(org.apache.calcite.rel.core.Join.class, filter.getInput()); + List expected = + joinType == JoinType.RIGHT ? Arrays.asList(null, 3) : Arrays.asList(1, null); + assertEquals(List.of(expected), rows(converted)); + } + + @Test + void postJoinFilterRunsBeforeEmit() { + Join join = equalityJoin(JoinType.LEFT); + Join filtered = + Join.builder() + .from(join) + .postJoinFilter(sb.isNull(sb.fieldReference(join, 1))) + .remap(sb.remap(0)) + .build(); + + assertEquals(List.of(List.of(1)), rows(substraitToCalcite.convert(filtered))); + } + + @ParameterizedTest + @EnumSource( + value = JoinType.class, + names = {"LEFT", "RIGHT", "OUTER"}) + void protoPostJoinFilterPreservesUnmatchedRows(JoinType joinType) { + Join join = equalityJoin(joinType); + Join filtered = + Join.builder() + .from(join) + .postJoinFilter( + sb.or(sb.isNull(sb.fieldReference(join, 0)), sb.isNull(sb.fieldReference(join, 1)))) + .remap(sb.remap(1, 0)) + .build(); + ExtensionCollector collector = new ExtensionCollector(); + io.substrait.proto.Rel proto = new RelProtoConverter(collector).toProto(filtered); + Rel decoded = new ProtoRelConverter(collector, extensions).from(proto); + + List> expected = + joinType == JoinType.LEFT + ? List.of(Arrays.asList(null, 1)) + : joinType == JoinType.RIGHT + ? List.of(Arrays.asList(3, null)) + : List.of(Arrays.asList(null, 1), Arrays.asList(3, null)); + assertEquals(expected, rows(substraitToCalcite.convert(decoded))); + } + + @ParameterizedTest + @EnumSource( + value = JoinType.class, + names = {"INNER", "LEFT", "LEFT_SEMI", "LEFT_ANTI"}) + void falseAndNullPostJoinFiltersProduceNoRows(JoinType joinType) { + Join join = equalityJoin(joinType); + for (Expression condition : List.of(sb.bool(false), ExpressionCreator.typedNull(N.BOOLEAN))) { + Join filtered = Join.builder().from(join).postJoinFilter(condition).build(); + + assertEquals(List.of(), rows(substraitToCalcite.convert(filtered))); + } + } + + @ParameterizedTest + @EnumSource( + value = JoinType.class, + names = {"LEFT_SEMI", "LEFT_ANTI"}) + void postJoinFilterUsesSemiAndAntiOutput(JoinType joinType) { + Join join = equalityJoin(joinType); + Join filtered = + Join.builder() + .from(join) + .postJoinFilter(sb.equal(sb.fieldReference(join, 0), sb.i32(1))) + .build(); + + assertEquals( + joinType == JoinType.LEFT_ANTI ? List.of(List.of(1)) : List.of(), + rows(substraitToCalcite.convert(filtered))); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void readFilterRetainsEnclosingCorrelation(boolean named) { + Rel outer = integers(1, 2).withRelAnchor(7); + FieldReference outerRef = + FieldReference.newRootStructOuterReferenceByRelReference(0, outer.getRecordType(), 7); + VirtualTableScan virtual = integers(1, 2); + Expression condition = sb.equal(sb.fieldReference(virtual, 0), outerRef); + Rel inner = + named + ? NamedScan.builder() + .initialSchema(virtual.getInitialSchema()) + .addNames("example") + .filter(condition) + .build() + : VirtualTableScan.builder().from(virtual).filter(condition).build(); + Rel plan = sb.filter(input -> sb.exists(inner), outer); + + Filter converted = assertInstanceOf(Filter.class, substraitToCalcite.convert(plan)); + assertFalse(converted.getVariablesSet().isEmpty()); + RexSubQuery exists = assertInstanceOf(RexSubQuery.class, converted.getCondition()); + Filter innerFilter = assertInstanceOf(Filter.class, exists.rel); + assertTrue(innerFilter.getCondition().toString().contains("$cor")); + } + + @Test + void postJoinFilterRetainsEnclosingCorrelation() { + Rel outer = integers(1, 2).withRelAnchor(7); + FieldReference outerRef = + FieldReference.newRootStructOuterReferenceByRelReference(0, outer.getRecordType(), 7); + Join join = equalityJoin(JoinType.INNER); + Join inner = + Join.builder() + .from(join) + .postJoinFilter(sb.equal(sb.fieldReference(join, 0), outerRef)) + .build(); + + Rel plan = sb.filter(input -> sb.exists(inner), outer); + Filter converted = assertInstanceOf(Filter.class, substraitToCalcite.convert(plan)); + assertFalse(converted.getVariablesSet().isEmpty()); + RexSubQuery exists = assertInstanceOf(RexSubQuery.class, converted.getCondition()); + Filter innerFilter = assertInstanceOf(Filter.class, exists.rel); + assertInstanceOf(org.apache.calcite.rel.core.Join.class, innerFilter.getInput()); + assertTrue(innerFilter.getCondition().toString().contains("$cor")); + } + + private Join equalityJoin(JoinType joinType) { + return sb.join( + input -> sb.equal(sb.fieldReference(input, 0), sb.fieldReference(input, 1)), + joinType, + integers(1, 2), + integers(2, 3)); + } + + private VirtualTableScan integers(int... values) { + return VirtualTableScan.builder() + .initialSchema(NamedStruct.of(List.of("id"), R.struct(R.I32))) + .rows( + Arrays.stream(values) + .mapToObj(value -> ExpressionCreator.nestedStruct(false, List.of(sb.i32(value)))) + .collect(Collectors.toList())) + .build(); + } + + private List> rows(RelNode rel) { + DataContext dataContext = + new DataContext() { + @Override + public SchemaPlus getRootSchema() { + return Frameworks.createRootSchema(true); + } + + @Override + public JavaTypeFactory getTypeFactory() { + return new JavaTypeFactoryImpl(); + } + + @Override + public QueryProvider getQueryProvider() { + return null; + } + + @Override + public Object get(String name) { + return null; + } + }; + RelNode executable = plan(rel, VirtualTableExpansionRule.instance()); + try (Interpreter interpreter = new Interpreter(dataContext, executable)) { + return interpreter.toList().stream().map(Arrays::asList).collect(Collectors.toList()); + } + } +} diff --git a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala index 712972859..8dcd55015 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala @@ -232,7 +232,7 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes throw new UnsupportedOperationException(s"Unsupported join type $other") } val plan = Join(left, right, joinType, condition, hint = JoinHint.NONE) - remap(plan, join.getRemap) + remap(applyFilter(plan, join.getPostJoinFilter, context), join.getRemap) } } @@ -424,7 +424,7 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes case _ => LocalRelation(ToSparkType.toAttributeSeq(virtualTableScan.getInitialSchema), rows) } - remap(plan, virtualTableScan.getRemap) + remap(applyFilter(plan, virtualTableScan.getFilter, context), virtualTableScan.getRemap) } override def visit( @@ -434,7 +434,7 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes case m: MultiInstanceRelation => m.newInstance() case other => other } - remap(plan, namedScan.getRemap) + remap(applyFilter(plan, namedScan.getFilter, context), namedScan.getRemap) } override def visit(localFiles: LocalFiles, context: EmptyVisitationContext): LogicalPlan = { @@ -467,7 +467,20 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes catalogTable = None, isStreaming = false ) - remap(plan, localFiles.getRemap) + remap(applyFilter(plan, localFiles.getFilter, context), localFiles.getRemap) + } + + private def applyFilter( + plan: LogicalPlan, + predicate: Optional[SExpression], + context: EmptyVisitationContext): LogicalPlan = { + if (predicate.isPresent) { + withChild(plan) { + Filter(predicate.get.accept(expressionConverter, context), plan) + } + } else { + plan + } } def convertFileFormat(fileFormat: FileFormat): (SparkFileFormat, Map[String, String]) = { diff --git a/spark/src/test/scala/io/substrait/spark/MandatoryPredicatesSuite.scala b/spark/src/test/scala/io/substrait/spark/MandatoryPredicatesSuite.scala new file mode 100644 index 000000000..45e402b55 --- /dev/null +++ b/spark/src/test/scala/io/substrait/spark/MandatoryPredicatesSuite.scala @@ -0,0 +1,188 @@ +package io.substrait.spark + +import io.substrait.spark.logical.{ToLogicalPlan, ToSubstraitRel} + +import org.apache.spark.SparkFunSuite +import org.apache.spark.sql.Row +import org.apache.spark.sql.classic.DatasetUtil +import org.apache.spark.sql.test.SharedSparkSession +import org.apache.spark.sql.types.StructType + +import io.substrait.`type`.TypeCreator +import io.substrait.dsl.SubstraitBuilder +import io.substrait.expression.{Expression, ExpressionCreator} +import io.substrait.relation.{AbstractReadRel, Join, LocalFiles => LocalFilesRel, NamedScan, Rel, VirtualTableScan} +import io.substrait.util.EmptyVisitationContext + +class MandatoryPredicatesSuite extends SparkFunSuite with SharedSparkSession { + + private val builder = new SubstraitBuilder + + override def beforeAll(): Unit = { + super.beforeAll() + sparkContext.setLogLevel("WARN") + } + + private def assertRows(rel: Rel, expected: Row*): Unit = { + val plan = rel.accept(new ToLogicalPlan(spark), EmptyVisitationContext.INSTANCE) + assert(plan.resolved) + val actual = DatasetUtil.fromLogicalPlan(spark, plan).collect().toSeq + assertResult(expected.sortBy(_.toString))(actual.sortBy(_.toString)) + } + + private def withScan(kind: String)(body: AbstractReadRel => Unit): Unit = { + val data = spark.sql("SELECT * FROM VALUES (1, true), (2, false), (3, NULL) AS t(id, keep)") + kind match { + case "named table" => + withTempView("predicate_scan") { + data.createOrReplaceTempView("predicate_scan") + body( + NamedScan + .builder() + .addNames("predicate_scan") + .initialSchema(ToSubstraitType.toNamedStruct(data.schema)) + .build()) + } + case "virtual table" => + body( + new ToSubstraitRel() + .visit(data.queryExecution.optimizedPlan) + .asInstanceOf[VirtualTableScan]) + case "local files" => + withTempPath { + path => + data.write.parquet(path.getAbsolutePath) + val read = spark.read.parquet(path.getAbsolutePath) + body( + new ToSubstraitRel() + .visit(read.queryExecution.optimizedPlan) + .asInstanceOf[LocalFilesRel]) + } + case other => throw new IllegalArgumentException(s"Unknown scan kind: $other") + } + } + + private def filteredScan(scan: AbstractReadRel, predicate: Expression): Rel = { + val remap = Rel.Remap.offset(0, 1) + scan match { + case named: NamedScan => + NamedScan.builder().from(named).filter(predicate).remap(remap).build() + case virtual: VirtualTableScan => + VirtualTableScan.builder().from(virtual).filter(predicate).remap(remap).build() + case files: LocalFilesRel => + LocalFilesRel.builder().from(files).filter(predicate).remap(remap).build() + case other => throw new IllegalArgumentException(s"Unknown scan: $other") + } + } + + Seq("named table", "virtual table", "local files").foreach { + kind => + test(s"mandatory read predicates on $kind run before emit") { + withScan(kind) { + scan => + assertRows(scan, Row(1, true), Row(2, false), Row(3, null)) + assertRows(filteredScan(scan, builder.bool(false))) + assertRows( + filteredScan(scan, ExpressionCreator.typedNull(TypeCreator.NULLABLE.BOOLEAN))) + // The predicate column is omitted from the output, and NULL must be rejected. + assertRows(filteredScan(scan, builder.fieldReference(scan, 1)), Row(1)) + } + } + } + + test("mandatory predicate on a zero-column virtual table") { + val scan = VirtualTableScan + .builder() + .initialSchema(ToSubstraitType.toNamedStruct(new StructType())) + .addRows(ExpressionCreator.nestedStruct(false)) + .build() + assertRows(scan, Row()) + assertRows(VirtualTableScan.builder().from(scan).filter(builder.bool(false)).build()) + } + + private def join(joinType: Join.JoinType): Join = { + val left = new ToSubstraitRel().visit( + spark.sql("SELECT * FROM VALUES (1), (2) AS l(id)").queryExecution.optimizedPlan) + val right = new ToSubstraitRel().visit( + spark.sql("SELECT * FROM VALUES (1), (3) AS r(id)").queryExecution.optimizedPlan) + builder.join( + (input: SubstraitBuilder.JoinInput) => + builder.equal(builder.fieldReference(input, 0), builder.fieldReference(input, 1)), + joinType, + left, + right) + } + + Seq(Join.JoinType.INNER, Join.JoinType.LEFT, Join.JoinType.RIGHT, Join.JoinType.OUTER).foreach { + joinType => + test(s"mandatory false post-join predicate on $joinType") { + val input = join(joinType) + assertRows(Join.builder().from(input).postJoinFilter(builder.bool(false)).build()) + } + } + + test("left outer post-join predicate sees null-extended output before emit") { + val input = join(Join.JoinType.LEFT) + assertRows(input, Row(1, 1), Row(2, null)) + val filtered = Join + .builder() + .from(input) + .postJoinFilter(builder.isNull(builder.fieldReference(input, 1))) + .remap(Rel.Remap.offset(0, 1)) + .build() + assertRows(filtered, Row(2)) + } + + test("right outer post-join predicate sees null-extended output before emit") { + val input = join(Join.JoinType.RIGHT) + assertRows(input, Row(1, 1), Row(null, 3)) + val filtered = Join + .builder() + .from(input) + .postJoinFilter(builder.isNull(builder.fieldReference(input, 0))) + .remap(Rel.Remap.offset(1, 1)) + .build() + assertRows(filtered, Row(3)) + } + + test("full outer post-join predicate filters both sides of the join") { + val input = join(Join.JoinType.OUTER) + val filtered = Join + .builder() + .from(input) + .postJoinFilter( + builder.or( + builder.isNull(builder.fieldReference(input, 0)), + builder.isNull(builder.fieldReference(input, 1)))) + .build() + assertRows(filtered, Row(2, null), Row(null, 3)) + } + + test("post-join predicate rejects null-extended rows") { + val input = join(Join.JoinType.LEFT) + val filtered = Join + .builder() + .from(input) + .postJoinFilter(builder.equal(builder.fieldReference(input, 1), builder.i32(1))) + .remap(Rel.Remap.offset(0, 1)) + .build() + assertRows(filtered, Row(1)) + } + + Seq(Join.JoinType.LEFT_SEMI, Join.JoinType.LEFT_ANTI).foreach { + joinType => + test(s"post-join predicate uses $joinType output") { + val input = join(joinType) + val filtered = Join + .builder() + .from(input) + .postJoinFilter(builder.equal(builder.fieldReference(input, 0), builder.i32(2))) + .build() + if (joinType == Join.JoinType.LEFT_ANTI) { + assertRows(filtered, Row(2)) + } else { + assertRows(filtered) + } + } + } +} From fa004ebdaef9c11cff3b39d1cca658976032e063 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Wed, 7 Oct 2026 02:00:00 -0400 Subject: [PATCH 2/2] fix: share predicate scope conversion and clarify coordinates --- .../substrait/relation/AbstractReadRel.java | 3 +- .../main/java/io/substrait/relation/Join.java | 4 +- .../io/substrait/relation/LateralJoin.java | 4 +- .../substrait/relation/physical/HashJoin.java | 4 +- .../relation/physical/MergeJoin.java | 4 +- .../type/proto/JoinRoundtripTest.java | 13 +---- .../isthmus/SubstraitRelNodeConverter.java | 32 ++++++++--- .../isthmus/EmbeddedPredicateTest.java | 53 ++++++++++++------- 8 files changed, 73 insertions(+), 44 deletions(-) diff --git a/core/src/main/java/io/substrait/relation/AbstractReadRel.java b/core/src/main/java/io/substrait/relation/AbstractReadRel.java index 6573a06c7..62def54a1 100644 --- a/core/src/main/java/io/substrait/relation/AbstractReadRel.java +++ b/core/src/main/java/io/substrait/relation/AbstractReadRel.java @@ -21,7 +21,8 @@ public abstract class AbstractReadRel extends ZeroInputRel implements HasExtensi public abstract NamedStruct getInitialSchema(); /** - * Returns an optional filter expression that must be applied during the read. + * Returns an optional filter expression that must be applied during the read. Its field + * references are indexed against the initial schema, before projection or emit remapping. * * @return the filter expression, if present */ diff --git a/core/src/main/java/io/substrait/relation/Join.java b/core/src/main/java/io/substrait/relation/Join.java index 1bef23c22..01cf1cb92 100644 --- a/core/src/main/java/io/substrait/relation/Join.java +++ b/core/src/main/java/io/substrait/relation/Join.java @@ -21,7 +21,9 @@ public abstract class Join extends BiRel implements HasExtension { public abstract Optional getCondition(); /** - * Returns the filter applied to the join output after the join is performed, if any. + * Returns the filter applied to the join output after the join is performed, if any. Its field + * references are indexed against the direct join output before emit remapping, including any + * null-extended fields produced by an outer join. * * @return the optional post-join filter */ diff --git a/core/src/main/java/io/substrait/relation/LateralJoin.java b/core/src/main/java/io/substrait/relation/LateralJoin.java index 6a8d69a4f..4c1487951 100644 --- a/core/src/main/java/io/substrait/relation/LateralJoin.java +++ b/core/src/main/java/io/substrait/relation/LateralJoin.java @@ -26,7 +26,9 @@ public abstract class LateralJoin extends BiRel implements HasExtension { public abstract Optional getCondition(); /** - * Returns the filter applied to the join output after the join is performed, if any. + * Returns the filter applied to the join output after the join is performed, if any. Its field + * references are indexed against the direct join output before emit remapping, including any + * null-extended fields produced by an outer join. * * @return the optional post-join filter */ diff --git a/core/src/main/java/io/substrait/relation/physical/HashJoin.java b/core/src/main/java/io/substrait/relation/physical/HashJoin.java index 11dab49c3..072f4860d 100644 --- a/core/src/main/java/io/substrait/relation/physical/HashJoin.java +++ b/core/src/main/java/io/substrait/relation/physical/HashJoin.java @@ -60,7 +60,9 @@ public List getRightKeys() { public abstract JoinType getJoinType(); /** - * Returns the filter applied to the join output after the join is performed, if any. + * Returns the filter applied to the join output after the join is performed, if any. Its field + * references are indexed against the direct join output before emit remapping, including any + * null-extended fields produced by an outer join. * * @return the optional post-join filter */ diff --git a/core/src/main/java/io/substrait/relation/physical/MergeJoin.java b/core/src/main/java/io/substrait/relation/physical/MergeJoin.java index 0ad957d6d..599079815 100644 --- a/core/src/main/java/io/substrait/relation/physical/MergeJoin.java +++ b/core/src/main/java/io/substrait/relation/physical/MergeJoin.java @@ -60,7 +60,9 @@ public List getRightKeys() { public abstract JoinType getJoinType(); /** - * Returns the filter applied to the join output after the join is performed, if any. + * Returns the filter applied to the join output after the join is performed, if any. Its field + * references are indexed against the direct join output before emit remapping, including any + * null-extended fields produced by an outer join. * * @return the optional post-join filter */ diff --git a/core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java b/core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java index d63ca42db..5dbdb8372 100644 --- a/core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java +++ b/core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java @@ -28,18 +28,7 @@ class JoinRoundtripTest extends TestBase { Arrays.asList(R.FP64, R.STRING, R.I64)); @ParameterizedTest - @EnumSource( - value = Join.JoinType.class, - names = { - "INNER", - "LEFT", - "RIGHT", - "OUTER", - "LEFT_SEMI", - "LEFT_ANTI", - "RIGHT_SEMI", - "RIGHT_ANTI" - }) + @EnumSource(value = Join.JoinType.class, names = "UNKNOWN", mode = EnumSource.Mode.EXCLUDE) void postJoinFilterUsesOutputSchemaBeforeEmit(Join.JoinType joinType) { Join join = sb.join( diff --git a/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java b/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java index 924cc86b5..9a592a81b 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java @@ -196,9 +196,9 @@ public static RelNode convert( @Override public RelNode visit(Filter filter, Context context) throws RuntimeException { RelNode input = filter.getInput().accept(this, context); - context.enterScope(AnchoredInput.of(filter.getInput().getRelAnchor(), input.getRowType())); - RexNode filterCondition = filter.getCondition().accept(expressionRexConverter, context); - RelNode node = relBuilder.push(input).filter(context.exitScope(), filterCondition).build(); + RelNode node = + applyFilter( + input, Optional.of(filter.getCondition()), context, filter.getInput().getRelAnchor()); return applyRelCommon(node, filter, input); } @@ -209,14 +209,30 @@ public RelNode visit(NamedScan namedScan, Context context) throws RuntimeExcepti return applyRelCommon(applyProjection(node, namedScan.getProjection()), namedScan); } - private RelNode applyFilter(RelNode input, Optional filter, Context context) { + /** + * Applies a filter against the input's direct row type, preserving correlation scopes. + * + * @param input the Calcite input before projection or emit remapping + * @param filter the optional filter to apply + * @param context the conversion context + * @return the input with the filter applied, or the input unchanged when no filter is present + */ + protected RelNode applyFilter(RelNode input, Optional filter, Context context) { + return applyFilter(input, filter, context, Optional.empty()); + } + + private RelNode applyFilter( + RelNode input, + Optional filter, + Context context, + Optional inputRelAnchor) { if (filter.isEmpty()) { return input; } - // Embedded predicates use the operator's direct row, before projection or emit mapping. This - // is an internal input, not another anchored Substrait relation; enclosing scopes still own - // any outer references used by the predicate. - context.enterScope(AnchoredInput.of(Optional.empty(), input.getRowType())); + // Predicates use the direct input row, before projection or emit mapping. Embedded filters on + // the relation being built have no anchor here; Filter relations retain their input anchor. + // Existing enclosing scopes remain available for outer references. + context.enterScope(AnchoredInput.of(inputRelAnchor, input.getRowType())); RexNode condition = filter.get().accept(expressionRexConverter, context); return relBuilder.push(input).filter(context.exitScope(), condition).build(); } diff --git a/isthmus/src/test/java/io/substrait/isthmus/EmbeddedPredicateTest.java b/isthmus/src/test/java/io/substrait/isthmus/EmbeddedPredicateTest.java index 37251a8af..630bd8316 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/EmbeddedPredicateTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/EmbeddedPredicateTest.java @@ -22,6 +22,7 @@ import java.util.Arrays; import java.util.List; import java.util.stream.Collectors; +import java.util.stream.Stream; import org.apache.calcite.DataContext; import org.apache.calcite.adapter.java.JavaTypeFactory; import org.apache.calcite.interpreter.Interpreter; @@ -34,11 +35,14 @@ import org.apache.calcite.rel.core.Values; import org.apache.calcite.rex.RexInputRef; import org.apache.calcite.rex.RexSubQuery; +import org.apache.calcite.rex.RexUtil; import org.apache.calcite.schema.SchemaPlus; import org.apache.calcite.tools.Frameworks; 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.EnumSource; +import org.junit.jupiter.params.provider.MethodSource; import org.junit.jupiter.params.provider.ValueSource; class EmbeddedPredicateTest extends PlanTestBase { @@ -123,16 +127,15 @@ void namedScanFiltersBeforeEmit() { assertRowMatch(project.getRowType(), R.I32); } - @Test - void namedScanFalseAndNullFiltersProduceNoRows() { + @ParameterizedTest(name = "mandatory named scan filter {0} produces no rows") + @MethodSource("falseAndNullPredicates") + void namedScanFalseAndNullFiltersProduceNoRows(String conditionName, Expression condition) { NamedScan scan = sb.namedScan(List.of("example"), List.of("id"), List.of(R.I32)); - for (Expression condition : List.of(sb.bool(false), ExpressionCreator.typedNull(N.BOOLEAN))) { - NamedScan filtered = NamedScan.builder().from(scan).filter(condition).build(); + NamedScan filtered = NamedScan.builder().from(scan).filter(condition).build(); - Values values = assertInstanceOf(Values.class, substraitToCalcite.convert(filtered)); - assertTrue(values.getTuples().isEmpty()); - assertRowMatch(values.getRowType(), R.I32); - } + Values values = assertInstanceOf(Values.class, substraitToCalcite.convert(filtered)); + assertEquals(List.of(), values.getTuples()); + assertRowMatch(values.getRowType(), R.I32); } @ParameterizedTest @@ -230,17 +233,14 @@ void protoPostJoinFilterPreservesUnmatchedRows(JoinType joinType) { assertEquals(expected, rows(substraitToCalcite.convert(decoded))); } - @ParameterizedTest - @EnumSource( - value = JoinType.class, - names = {"INNER", "LEFT", "LEFT_SEMI", "LEFT_ANTI"}) - void falseAndNullPostJoinFiltersProduceNoRows(JoinType joinType) { + @ParameterizedTest(name = "{0} post-join filter {1} produces no rows") + @MethodSource("falseAndNullPostJoinPredicates") + void falseAndNullPostJoinFiltersProduceNoRows( + JoinType joinType, String conditionName, Expression condition) { Join join = equalityJoin(joinType); - for (Expression condition : List.of(sb.bool(false), ExpressionCreator.typedNull(N.BOOLEAN))) { - Join filtered = Join.builder().from(join).postJoinFilter(condition).build(); + Join filtered = Join.builder().from(join).postJoinFilter(condition).build(); - assertEquals(List.of(), rows(substraitToCalcite.convert(filtered))); - } + assertEquals(List.of(), rows(substraitToCalcite.convert(filtered))); } @ParameterizedTest @@ -282,7 +282,7 @@ void readFilterRetainsEnclosingCorrelation(boolean named) { assertFalse(converted.getVariablesSet().isEmpty()); RexSubQuery exists = assertInstanceOf(RexSubQuery.class, converted.getCondition()); Filter innerFilter = assertInstanceOf(Filter.class, exists.rel); - assertTrue(innerFilter.getCondition().toString().contains("$cor")); + assertTrue(RexUtil.containsCorrelation(innerFilter.getCondition())); } @Test @@ -303,7 +303,22 @@ void postJoinFilterRetainsEnclosingCorrelation() { RexSubQuery exists = assertInstanceOf(RexSubQuery.class, converted.getCondition()); Filter innerFilter = assertInstanceOf(Filter.class, exists.rel); assertInstanceOf(org.apache.calcite.rel.core.Join.class, innerFilter.getInput()); - assertTrue(innerFilter.getCondition().toString().contains("$cor")); + assertTrue(RexUtil.containsCorrelation(innerFilter.getCondition())); + } + + static Stream falseAndNullPredicates() { + return Stream.of( + Arguments.of("FALSE", ExpressionCreator.bool(false, false)), + Arguments.of("NULL", ExpressionCreator.typedNull(N.BOOLEAN))); + } + + static Stream falseAndNullPostJoinPredicates() { + return Stream.of(JoinType.INNER, JoinType.LEFT, JoinType.LEFT_SEMI, JoinType.LEFT_ANTI) + .flatMap( + joinType -> + Stream.of( + Arguments.of(joinType, "FALSE", ExpressionCreator.bool(false, false)), + Arguments.of(joinType, "NULL", ExpressionCreator.typedNull(N.BOOLEAN)))); } private Join equalityJoin(JoinType joinType) {