Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
*/
Expand Down
4 changes: 3 additions & 1 deletion core/src/main/java/io/substrait/relation/Join.java
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,9 @@ public abstract class Join extends BiRel implements HasExtension {
public abstract Optional<Expression> 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
*/
Expand Down
4 changes: 3 additions & 1 deletion core/src/main/java/io/substrait/relation/LateralJoin.java
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,9 @@ public abstract class LateralJoin extends BiRel implements HasExtension {
public abstract Optional<Expression> 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
*/
Expand Down
39 changes: 26 additions & 13 deletions core/src/main/java/io/substrait/relation/ProtoRelConverter.java
Original file line number Diff line number Diff line change
Expand Up @@ -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()) {
Comment thread
nielspardon marked this conversation as resolved.
Comment thread
nielspardon marked this conversation as resolved.
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()));
Expand Down Expand Up @@ -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()));
}
Expand Down Expand Up @@ -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()));
}
Expand Down Expand Up @@ -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()));
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,9 @@ public List<FieldReference> 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
*/
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,9 @@ public List<FieldReference> 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
*/
Expand Down
79 changes: 70 additions & 9 deletions core/src/test/java/io/substrait/type/proto/JoinRoundtripTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -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 {

Expand All @@ -25,6 +27,28 @@ class JoinRoundtripTest extends TestBase {
Arrays.asList("d", "e", "f"),
Arrays.asList(R.FP64, R.STRING, R.I64));

@ParameterizedTest
@EnumSource(value = Join.JoinType.class, names = "UNKNOWN", mode = EnumSource.Mode.EXCLUDE)
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<Integer> leftKeys = Arrays.asList(0, 1);
Expand All @@ -36,6 +60,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<Integer> leftKeys = Arrays.asList(0, 1);
Expand Down Expand Up @@ -65,6 +104,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<Integer> leftKeys = Arrays.asList(0, 1);
Expand Down Expand Up @@ -130,24 +184,31 @@ 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)
.condition(
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());
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -196,18 +196,47 @@ 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);
}

@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);
}

/**
* 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<Expression> filter, Context context) {
return applyFilter(input, filter, context, Optional.empty());
}

private RelNode applyFilter(
RelNode input,
Optional<Expression> filter,
Context context,
Optional<Integer> inputRelAnchor) {
if (filter.isEmpty()) {
return input;
}
// 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();
}

@Override
public RelNode visit(LocalFiles localFiles, Context context) throws RuntimeException {
return visitFallback(localFiles, context);
Expand Down Expand Up @@ -257,6 +286,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);
}

Expand Down Expand Up @@ -1052,7 +1082,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 {
Expand All @@ -1063,7 +1096,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);
}
Expand Down
Loading
Loading