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
Original file line number Diff line number Diff line change
Expand Up @@ -277,13 +277,15 @@ public Plan visitLogicalJoin(LogicalJoin<? extends Plan, ? extends Plan> join, C
joinType = JoinType.INNER_JOIN;
}

return new LogicalJoin<>(joinType,
LogicalJoin<Plan, Plan> rewrittenJoin = new LogicalJoin<>(joinType,
newHashJoinConjuncts,
newOtherJoinConjuncts,
join.getMarkJoinConjuncts(),
join.getDistributeHint(),
join.getMarkJoinSlotReference(),
join.children(), join.getJoinReorderContext());
Plan eliminatedJoin = EliminateJoinCondition.eliminateJoinCondition(rewrittenJoin);
return eliminatedJoin == null ? rewrittenJoin : eliminatedJoin;
}

@Override
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,36 +19,89 @@

import org.apache.doris.nereids.rules.Rule;
import org.apache.doris.nereids.rules.RuleType;
import org.apache.doris.nereids.trees.expressions.Alias;
import org.apache.doris.nereids.trees.expressions.Expression;
import org.apache.doris.nereids.trees.expressions.NamedExpression;
import org.apache.doris.nereids.trees.expressions.Slot;
import org.apache.doris.nereids.trees.expressions.StatementScopeIdGenerator;
import org.apache.doris.nereids.trees.expressions.literal.BooleanLiteral;
import org.apache.doris.nereids.trees.expressions.literal.NullLiteral;
import org.apache.doris.nereids.trees.plans.Plan;
import org.apache.doris.nereids.trees.plans.logical.LogicalEmptyRelation;
import org.apache.doris.nereids.trees.plans.logical.LogicalJoin;
import org.apache.doris.nereids.trees.plans.logical.LogicalProject;

import com.google.common.collect.ImmutableList;

import java.util.List;
import java.util.Set;
import java.util.stream.Collectors;

/**
* Eliminate true Condition in Join Condition.
* Eliminate constant conditions in Join Condition.
*/
public class EliminateJoinCondition extends OneRewriteRuleFactory {

@Override
public Rule build() {
return logicalJoin().then(join -> {
List<Expression> hashJoinConjuncts = join.getHashJoinConjuncts().stream()
.filter(expression -> !expression.equals(BooleanLiteral.TRUE))
.collect(Collectors.toList());
List<Expression> otherJoinConjuncts = join.getOtherJoinConjuncts().stream()
.filter(expression -> !expression.equals(BooleanLiteral.TRUE))
.collect(Collectors.toList());
List<Expression> markJoinConjuncts = join.getMarkJoinConjuncts().stream()
.filter(expression -> !expression.equals(BooleanLiteral.TRUE))
.collect(Collectors.toList());
if (hashJoinConjuncts.size() == join.getHashJoinConjuncts().size()
&& otherJoinConjuncts.size() == join.getOtherJoinConjuncts().size()
&& markJoinConjuncts.size() == join.getMarkJoinConjuncts().size()) {
return null;
return logicalJoin()
.then(EliminateJoinCondition::eliminateJoinCondition)
.toRule(RuleType.ELIMINATE_JOIN_CONDITION);
}

static Plan eliminateJoinCondition(LogicalJoin<? extends Plan, ? extends Plan> join) {
List<Expression> hashJoinConjuncts = removeTrueConjuncts(join.getHashJoinConjuncts());
List<Expression> otherJoinConjuncts = removeTrueConjuncts(join.getOtherJoinConjuncts());
List<Expression> markJoinConjuncts = removeTrueConjuncts(join.getMarkJoinConjuncts());

if (!join.isMarkJoin() && (containsFalseOrNull(hashJoinConjuncts)
|| containsFalseOrNull(otherJoinConjuncts))) {
switch (join.getJoinType()) {
case INNER_JOIN:
case CROSS_JOIN:
return new LogicalEmptyRelation(StatementScopeIdGenerator.newRelationId(), join.getOutput());
case LEFT_OUTER_JOIN:
return projectNullPaddedJoinOutput(join, join.left());
case RIGHT_OUTER_JOIN:
return projectNullPaddedJoinOutput(join, join.right());
default:
break;
}
}

if (hashJoinConjuncts.size() == join.getHashJoinConjuncts().size()
&& otherJoinConjuncts.size() == join.getOtherJoinConjuncts().size()
&& markJoinConjuncts.size() == join.getMarkJoinConjuncts().size()) {
return null;
}
return join.withJoinConjuncts(hashJoinConjuncts, otherJoinConjuncts, markJoinConjuncts,
join.getJoinReorderContext());
}

private static List<Expression> removeTrueConjuncts(List<Expression> conjuncts) {
return conjuncts.stream()
.filter(expression -> !expression.equals(BooleanLiteral.TRUE))
.collect(Collectors.toList());
}

private static boolean containsFalseOrNull(List<Expression> conjuncts) {
return conjuncts.stream()
.anyMatch(expression -> expression.equals(BooleanLiteral.FALSE) || expression.isNullLiteral());
}

private static LogicalProject<Plan> projectNullPaddedJoinOutput(
LogicalJoin<? extends Plan, ? extends Plan> join, Plan preservedChild) {
Set<Slot> preservedOutput = preservedChild.getOutputSet();
ImmutableList.Builder<NamedExpression> projects =
ImmutableList.builderWithExpectedSize(join.getOutput().size());
for (Slot output : join.getOutput()) {
if (preservedOutput.contains(output)) {
projects.add(output);
} else {
projects.add(new Alias(output.getExprId(), ImmutableList.of(new NullLiteral(output.getDataType())),
output.getName(), output.getQualifier(), false));
}
return join.withJoinConjuncts(hashJoinConjuncts, otherJoinConjuncts, markJoinConjuncts,
join.getJoinReorderContext());
}).toRule(RuleType.ELIMINATE_JOIN_CONDITION);
}
return new LogicalProject<>(projects.build(), preservedChild);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -17,22 +17,34 @@

package org.apache.doris.nereids.rules.rewrite;

import org.apache.doris.nereids.trees.expressions.Alias;
import org.apache.doris.nereids.trees.expressions.EqualTo;
import org.apache.doris.nereids.trees.expressions.NamedExpression;
import org.apache.doris.nereids.trees.expressions.Slot;
import org.apache.doris.nereids.trees.expressions.functions.scalar.If;
import org.apache.doris.nereids.trees.expressions.literal.BooleanLiteral;
import org.apache.doris.nereids.trees.expressions.literal.NullLiteral;
import org.apache.doris.nereids.trees.plans.JoinType;
import org.apache.doris.nereids.trees.plans.logical.LogicalOlapScan;
import org.apache.doris.nereids.trees.plans.logical.LogicalPlan;
import org.apache.doris.nereids.trees.plans.logical.LogicalProject;
import org.apache.doris.nereids.util.LogicalPlanBuilder;
import org.apache.doris.nereids.util.MemoPatternMatchSupported;
import org.apache.doris.nereids.util.MemoTestUtils;
import org.apache.doris.nereids.util.PlanChecker;
import org.apache.doris.nereids.util.PlanConstructor;

import com.google.common.collect.ImmutableList;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.Test;

import java.util.List;
import java.util.Set;

class EliminateJoinConditionTest implements MemoPatternMatchSupported {
private final LogicalOlapScan scan1 = PlanConstructor.newLogicalOlapScan(0, "t1", 0);
private final LogicalOlapScan scan2 = PlanConstructor.newLogicalOlapScan(1, "t2", 0);
private final LogicalOlapScan scan3 = PlanConstructor.newLogicalOlapScan(2, "t3", 0);

@Test
void basicCase() {
Expand All @@ -48,4 +60,76 @@ void basicCase() {
&& join.getOtherJoinConjuncts().size() == 0)
);
}

@Test
void eliminateInnerJoinWithFalseCondition() {
LogicalPlan join = new LogicalPlanBuilder(scan1)
.join(scan2, JoinType.INNER_JOIN, ImmutableList.of(), ImmutableList.of(BooleanLiteral.FALSE))
.build();

PlanChecker.from(MemoTestUtils.createConnectContext(), join)
.applyTopDown(new EliminateJoinCondition())
.matches(logicalEmptyRelation());
}

@Test
void eliminateLeftOuterJoinWithNullCondition() {
LogicalPlan join = new LogicalPlanBuilder(scan1)
.join(scan2, JoinType.LEFT_OUTER_JOIN, ImmutableList.of(),
ImmutableList.of(NullLiteral.BOOLEAN_INSTANCE))
.build();

assertNullPaddedProject(join, scan1);
}

@Test
void eliminateRightOuterJoinWithFalseCondition() {
LogicalPlan join = new LogicalPlanBuilder(scan1)
.join(scan2, JoinType.RIGHT_OUTER_JOIN, ImmutableList.of(), ImmutableList.of(BooleanLiteral.FALSE))
.build();

assertNullPaddedProject(join, scan2);
}

private void assertNullPaddedProject(LogicalPlan join, LogicalPlan preservedChild) {
List<Slot> originalOutput = join.getOutput();
Set<Slot> preservedOutput = preservedChild.getOutputSet();

LogicalPlan rewritten = (LogicalPlan) PlanChecker.from(MemoTestUtils.createConnectContext(), join)
.applyTopDown(new EliminateJoinCondition())
.getPlan();
Assertions.assertInstanceOf(LogicalProject.class, rewritten);
LogicalProject<?> project = (LogicalProject<?>) rewritten;
Assertions.assertEquals(preservedChild, project.child());
Assertions.assertEquals(originalOutput, project.getOutput());
for (int i = 0; i < originalOutput.size(); i++) {
NamedExpression projectExpression = project.getProjects().get(i);
if (preservedOutput.contains(originalOutput.get(i))) {
Assertions.assertEquals(originalOutput.get(i), projectExpression);
} else {
Assertions.assertInstanceOf(Alias.class, projectExpression);
Assertions.assertInstanceOf(NullLiteral.class, projectExpression.child(0));
Assertions.assertEquals(originalOutput.get(i).getQualifier(), projectExpression.getQualifier());
}
}
}

@Test
void propagateNullPaddedOutputToInnerJoin() {
LogicalPlan leftOuterJoin = new LogicalPlanBuilder(scan1)
.join(scan2, JoinType.LEFT_OUTER_JOIN, ImmutableList.of(), ImmutableList.of(BooleanLiteral.FALSE))
.build();
Slot nullPaddedSlot = leftOuterJoin.getOutput().get(scan1.getOutput().size());
LogicalPlan innerJoin = new LogicalPlanBuilder(leftOuterJoin)
.join(scan3, JoinType.INNER_JOIN, ImmutableList.of(),
ImmutableList.of(new EqualTo(
new If(BooleanLiteral.TRUE, nullPaddedSlot, nullPaddedSlot),
scan3.getOutput().get(0))))
.build();

PlanChecker.from(MemoTestUtils.createConnectContext(), innerJoin)
.applyBottomUp(new EliminateJoinCondition())
.applyCustom(new ConstantPropagation())
.matches(logicalEmptyRelation());
}
}
Loading