Skip to content

Commit c437506

Browse files
committed
feat: stream aggregates over grouped input
1 parent c2fe167 commit c437506

1 file changed

Lines changed: 104 additions & 16 deletions

File tree

  • datafusion/physical-plan/src/aggregates

‎datafusion/physical-plan/src/aggregates/mod.rs‎

Lines changed: 104 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -1103,7 +1103,14 @@ impl AggregateExec {
11031103
input_order_mode = InputOrderMode::Linear;
11041104
}
11051105

1106-
let group_completion_mode = GroupCompletionMode::from(&input_order_mode);
1106+
let group_completion_mode = if !group_by.has_grouping_set()
1107+
&& !groupby_exprs.is_empty()
1108+
&& input_eq_properties.grouping_satisfy(groupby_exprs.iter().cloned())?
1109+
{
1110+
GroupCompletionMode::Full
1111+
} else {
1112+
GroupCompletionMode::from(&input_order_mode)
1113+
};
11071114

11081115
// construct a map from the input expression to the output expression of the Aggregation group by
11091116
let group_expr_mapping =
@@ -1122,6 +1129,16 @@ impl AggregateExec {
11221129
aggr_expr.as_ref(),
11231130
)?
11241131
};
1132+
// `compute_properties` derives emission from the public input-order
1133+
// mode. Override it only for the new case where unsorted input still
1134+
// has a group-completion guarantee.
1135+
let cache = if input_order_mode == InputOrderMode::Linear
1136+
&& group_completion_mode != GroupCompletionMode::None
1137+
{
1138+
cache.with_emission_type(input.pipeline_behavior())
1139+
} else {
1140+
cache
1141+
};
11251142

11261143
let mut exec = AggregateExec {
11271144
mode,
@@ -1494,6 +1511,10 @@ impl AggregateExec {
14941511
let mut eq_properties = input
14951512
.equivalence_properties()
14961513
.project(group_expr_mapping, schema);
1514+
// Grouping information is consumed by this aggregate. The aggregate
1515+
// output may have a different row layout, so do not pass explicit
1516+
// input grouping assertions to another aggregate.
1517+
eq_properties.clear_groupings();
14971518

14981519
// An aggregation that does not maintain its input order must not
14991520
// propegrate the input's ordering either, match `maintains_input_order` value
@@ -2454,7 +2475,7 @@ impl ExecutionPlan for AggregateExec {
24542475
required_input_ordering: _,
24552476
// Derived at construction from the input ordering and `group_by`.
24562477
input_order_mode: _,
2457-
// Derived at construction from `input_order_mode`.
2478+
// Derived at construction from the input properties and `group_by`.
24582479
group_completion_mode: _,
24592480
// Derived at construction by `Self::compute_properties`.
24602481
cache: _,
@@ -3364,7 +3385,7 @@ mod tests {
33643385
Int64Array, NullArray, StringArray, StructArray, UInt32Array, UInt64Array,
33653386
};
33663387
use arrow::compute::{SortOptions, concat_batches};
3367-
use arrow::datatypes::{Int32Type, Int64Type};
3388+
use arrow::datatypes::{Int32Type, Int64Type, TimeUnit};
33683389
use datafusion_common::test_util::{batches_to_sort_string, batches_to_string};
33693390
use datafusion_common::{DataFusionError, internal_err};
33703391
use datafusion_execution::config::SessionConfig;
@@ -3377,6 +3398,7 @@ mod tests {
33773398
Accumulator, AggregateUDF, AggregateUDFImpl, EmitTo, GroupsAccumulator, Operator,
33783399
Signature, Volatility,
33793400
};
3401+
use datafusion_functions::datetime::date_bin;
33803402
use datafusion_functions_aggregate::approx_percentile_cont::approx_percentile_cont_udaf;
33813403
use datafusion_functions_aggregate::array_agg::array_agg_udaf;
33823404
use datafusion_functions_aggregate::average::avg_udaf;
@@ -3385,10 +3407,9 @@ mod tests {
33853407
use datafusion_functions_aggregate::median::median_udaf;
33863408
use datafusion_functions_aggregate::min_max::min_udaf;
33873409
use datafusion_functions_aggregate::sum::sum_udaf;
3388-
use datafusion_physical_expr::Partitioning;
3389-
use datafusion_physical_expr::PhysicalSortExpr;
33903410
use datafusion_physical_expr::aggregate::AggregateExprBuilder;
33913411
use datafusion_physical_expr::expressions::{Literal, NotExpr, binary};
3412+
use datafusion_physical_expr::{Partitioning, PhysicalSortExpr, ScalarFunctionExpr};
33923413

33933414
use crate::projection::ProjectionExec;
33943415
use crate::repartition::RepartitionExec;
@@ -5551,7 +5572,7 @@ mod tests {
55515572
}
55525573

55535574
#[tokio::test]
5554-
async fn unsorted_contiguous_groups_use_final_emission() -> Result<()> {
5575+
async fn unsorted_contiguous_groups_use_incremental_emission() -> Result<()> {
55555576
let schema = Arc::new(Schema::new(vec![
55565577
Field::new("key", DataType::Int32, false),
55575578
Field::new("time_bin", DataType::Int64, false),
@@ -5579,18 +5600,21 @@ mod tests {
55795600
],
55805601
)?,
55815602
];
5603+
let key = col("key", &schema)?;
5604+
let time_bin = col("time_bin", &schema)?;
55825605
let group_by = PhysicalGroupBy::new_single(vec![
5583-
(col("key", &schema)?, "key".to_string()),
5584-
(col("time_bin", &schema)?, "time_bin".to_string()),
5606+
(Arc::clone(&key), "key".to_string()),
5607+
(Arc::clone(&time_bin), "time_bin".to_string()),
55855608
]);
55865609
let aggr_expr = Arc::new(
55875610
AggregateExprBuilder::new(sum_udaf(), vec![col("value", &schema)?])
55885611
.schema(Arc::clone(&schema))
55895612
.alias("SUM(value)")
55905613
.build()?,
55915614
);
5592-
let input: Arc<dyn ExecutionPlan> =
5593-
TestMemoryExec::try_new_exec(&[input_batches], Arc::clone(&schema), None)?;
5615+
let input = TestMemoryExec::try_new(&[input_batches], Arc::clone(&schema), None)?
5616+
.try_with_grouping_information(vec![vec![key, time_bin]])?;
5617+
let input: Arc<dyn ExecutionPlan> = Arc::new(input);
55945618
assert_eq!(input.output_partitioning().partition_count(), 1);
55955619

55965620
let aggregate = AggregateExec::try_new(
@@ -5603,15 +5627,20 @@ mod tests {
56035627
)?;
56045628

56055629
assert_eq!(aggregate.input_order_mode(), &InputOrderMode::Linear);
5606-
assert_eq!(aggregate.group_completion_mode, GroupCompletionMode::None);
5607-
// This captures the behavior before #24438. When the source can declare
5608-
// `(key, time_bin)` group-contiguous, the corresponding case can use
5609-
// `EmissionType::Incremental`.
5610-
assert_eq!(aggregate.cache().emission_type, EmissionType::Final);
5630+
assert_eq!(aggregate.group_completion_mode, GroupCompletionMode::Full);
5631+
assert_eq!(aggregate.cache().emission_type, EmissionType::Incremental);
5632+
assert!(
5633+
aggregate
5634+
.cache()
5635+
.equivalence_properties()
5636+
.geq_class()
5637+
.is_empty()
5638+
);
5639+
assert!(aggregate.cache().output_ordering().is_none());
56115640

56125641
let task_ctx = new_migrated_hash_ctx(1024);
56135642
let stream = aggregate.execute_typed(0, &task_ctx)?;
5614-
assert!(matches!(stream, StreamType::SingleHash(_)));
5643+
assert!(matches!(stream, StreamType::OrderedSingleAggregate(_)));
56155644
let stream: SendableRecordBatchStream = stream.into();
56165645
let output = collect(stream).await?;
56175646
assert_snapshot!(batches_to_sort_string(&output), @r"
@@ -5628,6 +5657,65 @@ mod tests {
56285657
Ok(())
56295658
}
56305659

5660+
#[test]
5661+
fn grouped_date_bin_projects_to_aggregate() -> Result<()> {
5662+
let schema = Arc::new(Schema::new(vec![
5663+
Field::new("key", DataType::Int32, false),
5664+
Field::new("time", DataType::Timestamp(TimeUnit::Second, None), false),
5665+
]));
5666+
let time_bin_expr = || -> Result<Arc<dyn PhysicalExpr>> {
5667+
Ok(Arc::new(ScalarFunctionExpr::try_new(
5668+
date_bin(),
5669+
vec![
5670+
lit(ScalarValue::new_interval_dt(0, 10_000)),
5671+
col("time", &schema)?,
5672+
],
5673+
&schema,
5674+
Arc::new(ConfigOptions::default()),
5675+
)?))
5676+
};
5677+
let input = TestMemoryExec::try_new(&[vec![]], Arc::clone(&schema), None)?
5678+
.try_with_grouping_information(vec![vec![
5679+
col("key", &schema)?,
5680+
time_bin_expr()?,
5681+
]])?;
5682+
let projection = ProjectionExec::try_new(
5683+
[
5684+
ProjectionExpr::new(col("key", &schema)?, "key"),
5685+
// Build this expression independently from the source
5686+
// assertion so the test exercises semantic expression matching.
5687+
ProjectionExpr::new(time_bin_expr()?, "time_bin"),
5688+
],
5689+
Arc::new(input),
5690+
)?;
5691+
5692+
let projected_schema = projection.schema();
5693+
let key = col("key", &projected_schema)?;
5694+
let time_bin = col("time_bin", &projected_schema)?;
5695+
assert!(
5696+
projection
5697+
.properties()
5698+
.equivalence_properties()
5699+
.grouping_satisfy([Arc::clone(&key), Arc::clone(&time_bin)])?
5700+
);
5701+
let aggregate = AggregateExec::try_new(
5702+
AggregateMode::Single,
5703+
PhysicalGroupBy::new_single(vec![
5704+
(key, "key".to_string()),
5705+
(time_bin, "time_bin".to_string()),
5706+
]),
5707+
vec![],
5708+
vec![],
5709+
Arc::new(projection),
5710+
projected_schema,
5711+
)?;
5712+
5713+
assert_eq!(aggregate.input_order_mode, InputOrderMode::Linear);
5714+
assert_eq!(aggregate.group_completion_mode, GroupCompletionMode::Full);
5715+
assert_eq!(aggregate.cache().emission_type, EmissionType::Incremental);
5716+
Ok(())
5717+
}
5718+
56315719
/// Ensures for ordered input, `OrderedPartialAggregateStream` is used.
56325720
#[tokio::test]
56335721
async fn ordered_partial_aggregate_planning() -> Result<()> {

0 commit comments

Comments
 (0)