This is an automated email from the ASF dual-hosted git repository. yiguolei pushed a commit to branch branch-4.2 in repository https://gitbox.apache.org/repos/asf/doris.git
commit ad01732816b09c18c111f15322e624791c8114c0 Author: morrySnow <[email protected]> AuthorDate: Wed Sep 16 23:30:56 2026 +0800 branch-4.1: [fix](aggregate) Normalize projected count slots before null safety checks #67732 (#68068) ### What problem does this PR solve? Related PR: #67732 Problem Summary: Backport the nullable projected COUNT safety fix to branch-4.1. Aggregate arguments are normalized through the Project before comparing them with filter slots, preventing unsafe COUNT_ON_INDEX pushdown for IS NULL predicates. The regression test is adapted to the 4.1 IndexType API. ### Release note None ### Check List (For Author) - Test - [x] Unit Test - Behavior changed: - [x] Yes. Unsafe storage-layer count pushdown is rejected for nullable projected count slots. - Does this need documentation? - [x] No. ### Check List (For Reviewer who merge this PR) - [ ] Confirm the release note - [ ] Confirm test cases - [ ] Confirm document - [ ] Add branch pick label --- .../rules/implementation/AggregateStrategies.java | 4 +-- .../rewrite/PhysicalStorageLayerAggregateTest.java | 38 ++++++++++++++++++++++ 2 files changed, 40 insertions(+), 2 deletions(-) diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/implementation/AggregateStrategies.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/implementation/AggregateStrategies.java index 5683ca528c9..9d4814aea80 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/implementation/AggregateStrategies.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/implementation/AggregateStrategies.java @@ -133,8 +133,8 @@ public class AggregateStrategies implements ImplementationRuleFactory { return false; } - Set<Slot> aggSlots = funcs.stream() - .flatMap(f -> f.getInputSlots().stream()) + Set<Slot> aggSlots = normalizeArguments(funcs, agg.child()).stream() + .flatMap(argument -> argument.getInputSlots().stream()) .collect(Collectors.toSet()); return aggSlots.isEmpty() || conjuncts.stream().allMatch(expr -> checkSlotInOrExpression(expr, aggSlots) && checkIsNullExpr(expr, aggSlots)); diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/PhysicalStorageLayerAggregateTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/PhysicalStorageLayerAggregateTest.java index 3c636d3018a..4846f3c004b 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/PhysicalStorageLayerAggregateTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/rewrite/PhysicalStorageLayerAggregateTest.java @@ -17,8 +17,10 @@ package org.apache.doris.nereids.rules.rewrite; +import org.apache.doris.analysis.IndexDef.IndexType; import org.apache.doris.catalog.Column; import org.apache.doris.catalog.DatabaseIf; +import org.apache.doris.catalog.Index; import org.apache.doris.catalog.TableIf; import org.apache.doris.catalog.Type; import org.apache.doris.datasource.CatalogIf; @@ -31,6 +33,7 @@ import org.apache.doris.nereids.rules.implementation.AggregateStrategies; import org.apache.doris.nereids.trees.expressions.Alias; import org.apache.doris.nereids.trees.expressions.Cast; import org.apache.doris.nereids.trees.expressions.Expression; +import org.apache.doris.nereids.trees.expressions.IsNull; import org.apache.doris.nereids.trees.expressions.functions.agg.Count; import org.apache.doris.nereids.trees.expressions.functions.agg.Max; import org.apache.doris.nereids.trees.expressions.functions.agg.Min; @@ -40,6 +43,7 @@ import org.apache.doris.nereids.trees.plans.RelationId; import org.apache.doris.nereids.trees.plans.logical.LogicalAggregate; import org.apache.doris.nereids.trees.plans.logical.LogicalFileScan; import org.apache.doris.nereids.trees.plans.logical.LogicalFileScan.SelectedPartitions; +import org.apache.doris.nereids.trees.plans.logical.LogicalFilter; import org.apache.doris.nereids.trees.plans.logical.LogicalOlapScan; import org.apache.doris.nereids.trees.plans.logical.LogicalProject; import org.apache.doris.nereids.trees.plans.logical.LogicalRelation; @@ -56,6 +60,7 @@ import org.apache.doris.nereids.util.PlanChecker; import org.apache.doris.nereids.util.PlanConstructor; import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableSet; import org.junit.jupiter.api.Test; import org.mockito.Mockito; @@ -368,6 +373,31 @@ public class PhysicalStorageLayerAggregateTest implements MemoPatternMatchSuppor ); } + @Test + public void testCountOnIndexRejectsIsNullOnProjectedCountSlot() { + LogicalOlapScan olapScan = PlanConstructor.newLogicalOlapScan(2, "count_alias", 0); + Index invertedIndex = new Index(1L, "idx_name", ImmutableList.of("name"), + IndexType.INVERTED, null, ""); + olapScan.getTable().getIndexIdToMeta().values().forEach( + meta -> meta.setIndexes(ImmutableList.of(invertedIndex))); + + LogicalFilter<LogicalOlapScan> filter = new LogicalFilter<>( + ImmutableSet.of(new IsNull(olapScan.getOutput().get(1))), olapScan); + LogicalProject<LogicalFilter<LogicalOlapScan>> project = new LogicalProject<>( + ImmutableList.of(new Alias(olapScan.getOutput().get(1), "x")), filter); + LogicalAggregate<LogicalProject<LogicalFilter<LogicalOlapScan>>> aggregate = new LogicalAggregate<>( + Collections.emptyList(), + ImmutableList.of(new Alias(new Count(project.getOutput().get(0)), "count_x"), + new Alias(new Count(), "count_star")), + true, Optional.empty(), project); + CascadesContext context = MemoTestUtils.createCascadesContext(aggregate); + context.getConnectContext().getSessionVariable().setEnablePushDownCountOnIndex(true); + + PlanChecker.from(context) + .applyImplementation(countOnIndex()) + .matches(logicalAggregate(logicalProject(logicalFilter(logicalOlapScan())))); + } + @Test void testProjectionCheck() { LogicalOlapScan olapScan = PlanConstructor.newLogicalOlapScan(1, "tbl", 0); @@ -418,4 +448,12 @@ public class PhysicalStorageLayerAggregateTest implements MemoPatternMatchSuppor .findFirst() .get(); } + + private Rule countOnIndex() { + return new AggregateStrategies().buildRules() + .stream() + .filter(rule -> rule.getRuleType() == RuleType.COUNT_ON_INDEX) + .findFirst() + .get(); + } } --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
