diff --git a/core/src/main/java/org/apache/calcite/plan/RelOptUtil.java b/core/src/main/java/org/apache/calcite/plan/RelOptUtil.java index b615c31d853..0e28d8a66d0 100644 --- a/core/src/main/java/org/apache/calcite/plan/RelOptUtil.java +++ b/core/src/main/java/org/apache/calcite/plan/RelOptUtil.java @@ -3402,6 +3402,44 @@ private static RexShuttle pushShuttle(final Calc calc) { }; } + /** + * Converts an expression that is based on the output fields of an + * {@link Aggregate} to an equivalent expression on the Aggregate's input + * fields. + * + *

An Aggregate's row type is {@code (group keys..., aggregate calls...)}, + * so output field {@code i} is input field {@code groupSet.nth(i)} for + * {@code i} less than {@link Aggregate#getGroupCount()}. Aggregate calls have + * no equivalent expression on the input, so {@code node} must reference only + * group keys; callers typically ensure this by classifying with + * {@link #splitFilters} against + * {@code ImmutableBitSet.range(aggregate.getGroupCount())}. + * + * @param node The expression to be converted + * @param aggregate Aggregate underneath the expression + * @return converted expression + */ + public static RexNode pushPastAggregate(RexNode node, Aggregate aggregate) { + return node.accept(pushShuttle(aggregate)); + } + + private static RexShuttle pushShuttle(final Aggregate aggregate) { + final List groupList = aggregate.getGroupSet().asList(); + final List inputFields = + aggregate.getInput().getRowType().getFieldList(); + return new RexShuttle() { + @Override public RexNode visitInputRef(RexInputRef ref) { + return RexInputRef.of(groupList.get(ref.getIndex()), inputFields); + } + + @Override public RexNode visitLambda(RexLambda lambda) { + // Lambda body references are at a different scope level. + // Do not remap indices inside lambda body against this aggregate. + return lambda; + } + }; + } + /** * Creates a new {@link org.apache.calcite.rel.rules.MultiJoin} to reflect * projection references from a diff --git a/core/src/main/java/org/apache/calcite/rel/metadata/RelMdDistinctRowCount.java b/core/src/main/java/org/apache/calcite/rel/metadata/RelMdDistinctRowCount.java index 7058dcbb4da..66a6763368b 100644 --- a/core/src/main/java/org/apache/calcite/rel/metadata/RelMdDistinctRowCount.java +++ b/core/src/main/java/org/apache/calcite/rel/metadata/RelMdDistinctRowCount.java @@ -190,6 +190,10 @@ protected RelMdDistinctRowCount() {} final RexBuilder rexBuilder = rel.getCluster().getRexBuilder(); RexNode childPreds = RexUtil.composeConjunction(rexBuilder, pushable, true); + // convert the predicate as it corresponds to the child input + if (childPreds != null) { + childPreds = RelOptUtil.pushPastAggregate(childPreds, rel); + } // set the bits as they correspond to the child input ImmutableBitSet.Builder childKey = ImmutableBitSet.builder(); diff --git a/core/src/main/java/org/apache/calcite/rel/metadata/RelMdSelectivity.java b/core/src/main/java/org/apache/calcite/rel/metadata/RelMdSelectivity.java index 5faaeaa85d3..fd07f9aedca 100644 --- a/core/src/main/java/org/apache/calcite/rel/metadata/RelMdSelectivity.java +++ b/core/src/main/java/org/apache/calcite/rel/metadata/RelMdSelectivity.java @@ -179,14 +179,20 @@ protected RelMdSelectivity() { @Nullable RexNode predicate) { final List notPushable = new ArrayList<>(); final List pushable = new ArrayList<>(); + // The predicate is expressed over the Aggregate's output, so only references + // to group keys (output fields below getGroupCount) can be pushed; and they + // must be converted to the input's field numbering before recursing. RelOptUtil.splitFilters( - rel.getGroupSet(), + ImmutableBitSet.range(rel.getGroupCount()), predicate, pushable, notPushable); final RexBuilder rexBuilder = rel.getCluster().getRexBuilder(); RexNode childPred = RexUtil.composeConjunction(rexBuilder, pushable, true); + if (childPred != null) { + childPred = RelOptUtil.pushPastAggregate(childPred, rel); + } Double selectivity = mq.getSelectivity(rel.getInput(), childPred); if (selectivity == null) { diff --git a/core/src/test/java/org/apache/calcite/test/RelMetadataTest.java b/core/src/test/java/org/apache/calcite/test/RelMetadataTest.java index 460e066051f..3dd1861eddb 100644 --- a/core/src/test/java/org/apache/calcite/test/RelMetadataTest.java +++ b/core/src/test/java/org/apache/calcite/test/RelMetadataTest.java @@ -101,6 +101,8 @@ import org.apache.calcite.rex.RexTableInputRef; import org.apache.calcite.rex.RexTableInputRef.RelTableRef; import org.apache.calcite.rex.RexUtil; +import org.apache.calcite.schema.SchemaPlus; +import org.apache.calcite.schema.impl.AbstractTable; import org.apache.calcite.sql.SqlBasicFunction; import org.apache.calcite.sql.SqlKind; import org.apache.calcite.sql.SqlOperator; @@ -111,6 +113,7 @@ import org.apache.calcite.sql.type.ReturnTypes; import org.apache.calcite.sql.type.SqlTypeName; import org.apache.calcite.test.catalog.MockCatalogReaderSimple; +import org.apache.calcite.tools.FrameworkConfig; import org.apache.calcite.tools.Frameworks; import org.apache.calcite.tools.RelBuilder; import org.apache.calcite.util.ArrowSet; @@ -1831,6 +1834,154 @@ void testColumnOriginsUnion() { isAlmost(DEFAULT_COMP_SELECTIVITY * DEFAULT_EQUAL_SELECTIVITY)); } + /** Null fraction of each column of the table used by the + * {@code testSelectivityAggregate} and + * {@code testDistinctRowCountAggregate} tests; distinct per column, so a + * selectivity identifies the column that the predicate resolved to. */ + private static final double[] NULL_FRACTION = {0.13, 0.42, 0.77}; // a, b, c + + /** Test case for + * [CALCITE-7687] + * RelMdSelectivity and RelMdDistinctRowCount for Aggregate can propagate a + * predicate with wrong references. + * + *

An {@link Aggregate}'s output field {@code i} is input field + * {@code groupSet.nth(i)}, so a predicate on a group key must be converted + * before it is pushed to the input. */ + @Test void testSelectivityAggregateConvertsPredicateToInputFields() { + final SelectivityByColumnTable table = new SelectivityByColumnTable(); + final RelNode agg = aggregateGroupingOnFields1And2(table); + final RelMetadataQuery mq = agg.getCluster().getMetadataQuery(); + + // Aggregate output $1 is column "c", input $2. + assertThat(mq.getSelectivity(agg, isNullOn(agg, 1)), isAlmost(NULL_FRACTION[2])); + assertThat(table.received, hasToString("IS NULL($2)")); + + // Aggregate output $0 is column "b", input $1. + assertThat(mq.getSelectivity(agg, isNullOn(agg, 0)), isAlmost(NULL_FRACTION[1])); + assertThat(table.received, hasToString("IS NULL($1)")); + } + + /** Test case for + * [CALCITE-7687] + * RelMdSelectivity and RelMdDistinctRowCount for Aggregate can propagate a + * predicate with wrong references. + * + *

A predicate on an aggregate call must not be pushed to the + * {@link Aggregate}'s input, because an aggregate call has no equivalent + * expression there. */ + @Test void testSelectivityAggregateDoesNotPushAggregateCallPredicate() { + final SelectivityByColumnTable table = new SelectivityByColumnTable(); + final RelNode agg = aggregateGroupingOnFields1And2(table); + final RelMetadataQuery mq = agg.getCluster().getMetadataQuery(); + + // Aggregate output $2 is COUNT($0), not a group key. + assertThat(mq.getSelectivity(agg, isNullOn(agg, 2)), isAlmost(DEFAULT_SELECTIVITY)); + assertThat(table.received, nullValue()); + } + + /** Test case for + * [CALCITE-7687] + * RelMdSelectivity and RelMdDistinctRowCount for Aggregate can propagate a + * predicate with wrong references. + * + *

As {@link #testSelectivityAggregateConvertsPredicateToInputFields()}, + * but for {@code getDistinctRowCount}: the group key was already converted by + * {@link RelMdUtil#setAggChildKeys}, the predicate beside it was not. */ + @Test void testDistinctRowCountAggregateConvertsPredicateToInputFields() { + final DistinctRowCountByColumnTable table = new DistinctRowCountByColumnTable(); + final RelNode agg = aggregateGroupingOnFields1And2(table); + final RelMetadataQuery mq = agg.getCluster().getMetadataQuery(); + + // Group key is output $0 ("b", input $1); predicate is on output $1 ("c", input $2). + mq.getDistinctRowCount(agg, ImmutableBitSet.of(0), isNullOn(agg, 1)); + assertThat(table.receivedGroupKey, hasToString("{1}")); + assertThat(table.receivedPredicate, hasToString("IS NULL($2)")); + } + + /** Returns {@code Aggregate(group={1, 2}, COUNT($0))} over a scan of + * {@code table}, whose row type is {@code (a, b, c)}; the group set is + * deliberately not an identity prefix. */ + private static RelNode aggregateGroupingOnFields1And2(AbstractTable table) { + final SchemaPlus root = Frameworks.createRootSchema(true); + root.add("T", table); + final FrameworkConfig config = + Frameworks.newConfigBuilder().defaultSchema(root).build(); + final RelBuilder b = RelBuilder.create(config); + return b.scan("T") + .aggregate(b.groupKey(1, 2), b.count(false, "cnt", b.field(0))) + .build(); + } + + private static RexNode isNullOn(RelNode rel, int i) { + final RexBuilder rexBuilder = rel.getCluster().getRexBuilder(); + return rexBuilder.makeCall(SqlStdOperatorTable.IS_NULL, + rexBuilder.makeInputRef(rel.getRowType().getFieldList().get(i).getType(), i)); + } + + private static RelDataType abcRowType(RelDataTypeFactory typeFactory) { + final RelDataType varchar = + typeFactory.createTypeWithNullability( + typeFactory.createSqlType(SqlTypeName.VARCHAR), true); + return typeFactory.builder() + .add("a", varchar).add("b", varchar).add("c", varchar).build(); + } + + /** Returns the selectivity implied by the null fraction of the column that + * {@code predicate} names, or Calcite's generic guess if it names none. */ + private static @Nullable Double selectivityOf(@Nullable RexNode predicate) { + if (predicate instanceof RexCall + && (predicate.getKind() == SqlKind.IS_NULL + || predicate.getKind() == SqlKind.IS_NOT_NULL) + && ((RexCall) predicate).getOperands().get(0) instanceof RexInputRef) { + final RexInputRef ref = + (RexInputRef) ((RexCall) predicate).getOperands().get(0); + final double nullFraction = NULL_FRACTION[ref.getIndex()]; + return predicate.getKind() == SqlKind.IS_NULL + ? nullFraction + : 1.0 - nullFraction; + } + return RelMdUtil.guessSelectivity(predicate); + } + + /** Table that derives selectivity from per-column null fractions, and records + * the predicate it was given. Reached via + * {@link org.apache.calcite.plan.RelOptTable#unwrap}. */ + private static class SelectivityByColumnTable extends AbstractTable + implements BuiltInMetadata.Selectivity.Handler { + @Nullable RexNode received; + + @Override public RelDataType getRowType(RelDataTypeFactory typeFactory) { + return abcRowType(typeFactory); + } + + @Override public @Nullable Double getSelectivity(RelNode r, + RelMetadataQuery mq, @Nullable RexNode predicate) { + received = predicate; + return selectivityOf(predicate); + } + } + + /** Table that records the group key and predicate that + * {@code getDistinctRowCount} passes down. */ + private static class DistinctRowCountByColumnTable extends AbstractTable + implements BuiltInMetadata.DistinctRowCount.Handler { + @Nullable RexNode receivedPredicate; + @Nullable ImmutableBitSet receivedGroupKey; + + @Override public RelDataType getRowType(RelDataTypeFactory typeFactory) { + return abcRowType(typeFactory); + } + + @Override public @Nullable Double getDistinctRowCount(RelNode r, + RelMetadataQuery mq, ImmutableBitSet groupKey, + @Nullable RexNode predicate) { + receivedGroupKey = groupKey; + receivedPredicate = predicate; + return 1000.0; + } + } + /** Test case for * [CALCITE-1808] * JaninoRelMetadataProvider loading cache might cause