Skip to content
Draft
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
13 changes: 9 additions & 4 deletions core/src/main/java/io/substrait/relation/ProtoRelConverter.java
Original file line number Diff line number Diff line change
Expand Up @@ -1013,15 +1013,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()));
Expand Down
35 changes: 35 additions & 0 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,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<Integer> leftKeys = Arrays.asList(0, 1);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -204,9 +204,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(node, namedScan);
}

private RelNode applyFilter(RelNode input, Optional<Expression> filter, Context context) {
if (filter.isEmpty()) {
return input;
}
// Embedded predicates use the operator's direct row, before its 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);
Expand Down Expand Up @@ -256,6 +269,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 @@ -921,7 +935,10 @@ public RelNode visit(VirtualTableScan virtualTableScan, Context context) {
tuplesBuilder.add(tupleBuilder.build());
}
return applyRelCommon(
LogicalValues.create(relBuilder.getCluster(), rowType, tuplesBuilder.build()),
applyFilter(
LogicalValues.create(relBuilder.getCluster(), rowType, tuplesBuilder.build()),
virtualTableScan.getFilter(),
context),
virtualTableScan);
} else {
// A row that does not fit a LogicalValues tuple keeps its expressions, in a relation of our
Expand All @@ -930,7 +947,11 @@ public RelNode visit(VirtualTableScan virtualTableScan, Context context) {
// consumer whose planner only knows Calcite's own relations can expand it with
// VirtualTableExpansionRule.
return applyRelCommon(
VirtualTable.create(relBuilder.getCluster(), rowType, convertedRows), virtualTableScan);
applyFilter(
VirtualTable.create(relBuilder.getCluster(), rowType, convertedRows),
virtualTableScan.getFilter(),
context),
virtualTableScan);
}
}

Expand Down
Loading
Loading