diff --git a/datafusion/functions-aggregate/src/min_max.rs b/datafusion/functions-aggregate/src/min_max.rs index 89a1e8114f5e4..d3b1d327e19dc 100644 --- a/datafusion/functions-aggregate/src/min_max.rs +++ b/datafusion/functions-aggregate/src/min_max.rs @@ -46,8 +46,9 @@ use crate::min_max::min_max_bytes::MinMaxBytesAccumulator; use crate::min_max::min_max_struct::MinMaxStructAccumulator; use datafusion_common::ScalarValue; use datafusion_expr::{ - Accumulator, AggregateUDFImpl, Documentation, SetMonotonicity, Signature, Volatility, - function::AccumulatorArgs, + Accumulator, AggregateUDFImpl, Documentation, Expr, SetMonotonicity, Signature, + Volatility, + function::{AccumulatorArgs, AggregateFunctionSimplification}, }; use datafusion_expr::{GroupsAccumulator, StatisticsArgs}; use datafusion_macros::user_doc; @@ -685,6 +686,14 @@ impl AggregateUDFImpl for Min { datafusion_expr::ReversedUDAF::Identical } + fn simplify(&self) -> Option { + // `min(DISTINCT x)` is identical to `min(x)`, therefore drop DISTINCT. + Some(Box::new(|mut aggregate_function, _info| { + aggregate_function.params.distinct = false; + Ok(Expr::AggregateFunction(aggregate_function)) + })) + } + fn documentation(&self) -> Option<&Documentation> { self.doc() } diff --git a/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs b/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs index 5436bd092163e..19fbfb62fcd15 100644 --- a/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs +++ b/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs @@ -2489,6 +2489,7 @@ mod tests { interval_arithmetic::Interval, *, }; + use datafusion_functions_aggregate::min_max::min_udaf; use datafusion_functions_window_common::field::WindowUDFFieldArgs; use datafusion_functions_window_common::partition::PartitionEvaluatorArgs; use datafusion_physical_expr::PhysicalExpr; @@ -5395,6 +5396,23 @@ mod tests { assert_eq!(simplify(aggregate_function_expr), expected); } + #[test] + fn test_simplify_min_drops_distinct() { + let min_agg = |distinct: bool| { + Expr::AggregateFunction(expr::AggregateFunction::new_udf( + min_udaf(), + vec![col("c3")], + distinct, + None, + vec![], + None, + )) + }; + + let simplified = simplify(min_agg(true)); + assert_eq!(simplified, min_agg(false)); + } + /// A Mock UDAF which defines `simplify` to be used in tests /// related to UDAF simplification #[derive(Debug, Clone, PartialEq, Eq, Hash)] diff --git a/datafusion/optimizer/src/single_distinct_to_groupby.rs b/datafusion/optimizer/src/single_distinct_to_groupby.rs index 00c8fab228117..88aefe17d6ba6 100644 --- a/datafusion/optimizer/src/single_distinct_to_groupby.rs +++ b/datafusion/optimizer/src/single_distinct_to_groupby.rs @@ -61,6 +61,15 @@ impl SingleDistinctToGroupBy { } } +fn unalias_aggregate(expr: &Expr) -> &Expr { + match expr { + Expr::Alias(alias) if matches!(*alias.expr, Expr::AggregateFunction(_)) => { + &alias.expr + } + _ => expr, + } +} + /// Check whether all aggregate exprs are distinct on a single field. fn is_single_distinct_agg(aggr_expr: &[Expr]) -> Result { let mut fields_set = HashSet::new(); @@ -76,7 +85,7 @@ fn is_single_distinct_agg(aggr_expr: &[Expr]) -> Result { order_by, null_treatment: _, }, - }) = expr + }) = unalias_aggregate(expr) { if filter.is_some() || !order_by.is_empty() { return Ok(false); @@ -179,6 +188,7 @@ impl OptimizerRule for SingleDistinctToGroupBy { let mut inner_aggr_exprs = vec![]; let outer_aggr_exprs = aggr_expr .into_iter() + .map(|aggr_expr| aggr_expr.unalias()) .map(|aggr_expr| match aggr_expr { Expr::AggregateFunction(AggregateFunction { func, diff --git a/datafusion/sqllogictest/test_files/group_by.slt b/datafusion/sqllogictest/test_files/group_by.slt index 08d2ee509f192..332f8190c61e5 100644 --- a/datafusion/sqllogictest/test_files/group_by.slt +++ b/datafusion/sqllogictest/test_files/group_by.slt @@ -4249,6 +4249,29 @@ physical_plan 07)------------AggregateExec: mode=Partial, gby=[y@1 as y, CAST(x@0 AS Float64) as alias1], aggr=[] 08)--------------DataSourceExec: partitions=1, partition_sizes=[1] + +statement ok +CREATE TABLE min_distinct(g int, x int) AS VALUES + (1, 3), (1, 3), (1, 1), (1, NULL), + (2, NULL), (2, NULL), + (3, 7), (3, 5), (3, 5); + +query TT +EXPLAIN SELECT g, min(DISTINCT x) FROM min_distinct GROUP BY g; +---- +logical_plan +01)Aggregate: groupBy=[[min_distinct.g]], aggr=[[min(min_distinct.x) AS min(DISTINCT min_distinct.x)]] +02)--TableScan: min_distinct projection=[g, x] +physical_plan +01)AggregateExec: mode=FinalPartitioned, gby=[g@0 as g], aggr=[min(min_distinct.x) as min(DISTINCT min_distinct.x)] +02)--RepartitionExec: partitioning=Hash([g@0], 8), input_partitions=8 +03)----AggregateExec: mode=Partial, gby=[g@0 as g], aggr=[min(min_distinct.x) as min(DISTINCT min_distinct.x)] +04)------RepartitionExec: partitioning=RoundRobinBatch(8), input_partitions=1 +05)--------DataSourceExec: partitions=1, partition_sizes=[5] + +statement ok +DROP TABLE min_distinct; + # create an unbounded table that contains ordered timestamp. statement ok CREATE UNBOUNDED EXTERNAL TABLE unbounded_csv_with_timestamps ( @@ -4432,20 +4455,20 @@ EXPLAIN SELECT c1, count(distinct c2), min(distinct c2), sum(c3), max(c4) FROM a ---- logical_plan 01)Sort: aggregate_test_100.c1 ASC NULLS LAST -02)--Projection: aggregate_test_100.c1, count(alias1) AS count(DISTINCT aggregate_test_100.c2), min(alias1) AS min(DISTINCT aggregate_test_100.c2), sum(alias2) AS sum(aggregate_test_100.c3), max(alias3) AS max(aggregate_test_100.c4) -03)----Aggregate: groupBy=[[aggregate_test_100.c1]], aggr=[[count(alias1), min(alias1), sum(alias2), max(alias3)]] -04)------Aggregate: groupBy=[[aggregate_test_100.c1, aggregate_test_100.c2 AS alias1]], aggr=[[sum(CAST(aggregate_test_100.c3 AS Int64)) AS alias2, max(aggregate_test_100.c4) AS alias3]] +02)--Projection: aggregate_test_100.c1, count(alias1) AS count(DISTINCT aggregate_test_100.c2), min(alias2) AS min(DISTINCT aggregate_test_100.c2), sum(alias3) AS sum(aggregate_test_100.c3), max(alias4) AS max(aggregate_test_100.c4) +03)----Aggregate: groupBy=[[aggregate_test_100.c1]], aggr=[[count(alias1), min(alias2), sum(alias3), max(alias4)]] +04)------Aggregate: groupBy=[[aggregate_test_100.c1, aggregate_test_100.c2 AS alias1]], aggr=[[min(aggregate_test_100.c2) AS alias2, sum(CAST(aggregate_test_100.c3 AS Int64)) AS alias3, max(aggregate_test_100.c4) AS alias4]] 05)--------TableScan: aggregate_test_100 projection=[c1, c2, c3, c4] physical_plan 01)SortPreservingMergeExec: [c1@0 ASC NULLS LAST] -02)--ProjectionExec: expr=[c1@0 as c1, count(alias1)@1 as count(DISTINCT aggregate_test_100.c2), min(alias1)@2 as min(DISTINCT aggregate_test_100.c2), sum(alias2)@3 as sum(aggregate_test_100.c3), max(alias3)@4 as max(aggregate_test_100.c4)] +02)--ProjectionExec: expr=[c1@0 as c1, count(alias1)@1 as count(DISTINCT aggregate_test_100.c2), min(alias2)@2 as min(DISTINCT aggregate_test_100.c2), sum(alias3)@3 as sum(aggregate_test_100.c3), max(alias4)@4 as max(aggregate_test_100.c4)] 03)----SortExec: expr=[c1@0 ASC NULLS LAST], preserve_partitioning=[true] -04)------AggregateExec: mode=FinalPartitioned, gby=[c1@0 as c1], aggr=[count(alias1), min(alias1), sum(alias2), max(alias3)] +04)------AggregateExec: mode=FinalPartitioned, gby=[c1@0 as c1], aggr=[count(alias1), min(alias2), sum(alias3), max(alias4)] 05)--------RepartitionExec: partitioning=Hash([c1@0], 8), input_partitions=8 -06)----------AggregateExec: mode=Partial, gby=[c1@0 as c1], aggr=[count(alias1), min(alias1), sum(alias2), max(alias3)] -07)------------AggregateExec: mode=FinalPartitioned, gby=[c1@0 as c1, alias1@1 as alias1], aggr=[sum(aggregate_test_100.c3) as alias2, max(aggregate_test_100.c4) as alias3] +06)----------AggregateExec: mode=Partial, gby=[c1@0 as c1], aggr=[count(alias1), min(alias2), sum(alias3), max(alias4)] +07)------------AggregateExec: mode=FinalPartitioned, gby=[c1@0 as c1, alias1@1 as alias1], aggr=[min(aggregate_test_100.c2) as alias2, sum(aggregate_test_100.c3) as alias3, max(aggregate_test_100.c4) as alias4] 08)--------------RepartitionExec: partitioning=Hash([c1@0, alias1@1], 8), input_partitions=8 -09)----------------AggregateExec: mode=Partial, gby=[c1@0 as c1, c2@1 as alias1], aggr=[sum(aggregate_test_100.c3) as alias2, max(aggregate_test_100.c4) as alias3] +09)----------------AggregateExec: mode=Partial, gby=[c1@0 as c1, c2@1 as alias1], aggr=[min(aggregate_test_100.c2) as alias2, sum(aggregate_test_100.c3) as alias3, max(aggregate_test_100.c4) as alias4] 10)------------------RepartitionExec: partitioning=RoundRobinBatch(8), input_partitions=1 11)--------------------DataSourceExec: file_groups={1 group: [[WORKSPACE_ROOT/testing/data/csv/aggregate_test_100.csv]]}, projection=[c1, c2, c3, c4], file_type=csv, has_header=true