diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/ConstantPropagation.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/ConstantPropagation.java index 1d77eb02b7341c..12570f9503533f 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/ConstantPropagation.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/ConstantPropagation.java @@ -277,13 +277,15 @@ public Plan visitLogicalJoin(LogicalJoin join, C joinType = JoinType.INNER_JOIN; } - return new LogicalJoin<>(joinType, + LogicalJoin 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 diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/EliminateJoinCondition.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/EliminateJoinCondition.java index 23625afb583f8b..82ec4bc7b81e42 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/EliminateJoinCondition.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/EliminateJoinCondition.java @@ -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 hashJoinConjuncts = join.getHashJoinConjuncts().stream() - .filter(expression -> !expression.equals(BooleanLiteral.TRUE)) - .collect(Collectors.toList()); - List otherJoinConjuncts = join.getOtherJoinConjuncts().stream() - .filter(expression -> !expression.equals(BooleanLiteral.TRUE)) - .collect(Collectors.toList()); - List 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 join) { + List hashJoinConjuncts = removeTrueConjuncts(join.getHashJoinConjuncts()); + List otherJoinConjuncts = removeTrueConjuncts(join.getOtherJoinConjuncts()); + List 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 removeTrueConjuncts(List conjuncts) { + return conjuncts.stream() + .filter(expression -> !expression.equals(BooleanLiteral.TRUE)) + .collect(Collectors.toList()); + } + + private static boolean containsFalseOrNull(List conjuncts) { + return conjuncts.stream() + .anyMatch(expression -> expression.equals(BooleanLiteral.FALSE) || expression.isNullLiteral()); + } + + private static LogicalProject projectNullPaddedJoinOutput( + LogicalJoin join, Plan preservedChild) { + Set preservedOutput = preservedChild.getOutputSet(); + ImmutableList.Builder 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); } } diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/EliminateJoinConditionTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/EliminateJoinConditionTest.java index 37acd78e027345..c34942b7b5ff7c 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/EliminateJoinConditionTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/EliminateJoinConditionTest.java @@ -17,10 +17,17 @@ 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; @@ -28,11 +35,16 @@ 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() { @@ -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 originalOutput = join.getOutput(); + Set 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()); + } }