use datafusion_physical_expr::{
AggregateExpr, EmitTo, GroupsAccumulator, GroupsAccumulatorAdapter,
};
use log::debug;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::vec;
use futures::ready;
use futures::stream::{Stream, StreamExt};
use crate::physical_plan::aggregates::group_values::{new_group_values, GroupValues};
use crate::physical_plan::aggregates::{
evaluate_group_by, evaluate_many, evaluate_optional, group_schema, AggregateMode,
PhysicalGroupBy,
};
use crate::physical_plan::metrics::{BaselineMetrics, RecordOutput};
use crate::physical_plan::{aggregates, PhysicalExpr};
use crate::physical_plan::{RecordBatchStream, SendableRecordBatchStream};
use arrow::array::*;
use arrow::{datatypes::SchemaRef, record_batch::RecordBatch};
use datafusion_common::Result;
use datafusion_execution::memory_pool::proxy::VecAllocExt;
use datafusion_execution::memory_pool::{MemoryConsumer, MemoryReservation};
use datafusion_execution::TaskContext;
#[derive(Debug, Clone)]
pub(crate) enum ExecutionState {
ReadingInput,
ProducingOutput(RecordBatch),
Done,
}
use super::order::GroupOrdering;
use super::AggregateExec;
pub(crate) struct GroupedHashAggregateStream {
schema: SchemaRef,
input: SendableRecordBatchStream,
mode: AggregateMode,
accumulators: Vec<Box<dyn GroupsAccumulator>>,
aggregate_arguments: Vec<Vec<Arc<dyn PhysicalExpr>>>,
filter_expressions: Vec<Option<Arc<dyn PhysicalExpr>>>,
group_by: PhysicalGroupBy,
reservation: MemoryReservation,
group_values: Box<dyn GroupValues>,
current_group_indices: Vec<usize>,
exec_state: ExecutionState,
baseline_metrics: BaselineMetrics,
batch_size: usize,
group_ordering: GroupOrdering,
input_done: bool,
}
impl GroupedHashAggregateStream {
pub fn new(
agg: &AggregateExec,
context: Arc<TaskContext>,
partition: usize,
) -> Result<Self> {
debug!("Creating GroupedHashAggregateStream");
let agg_schema = Arc::clone(&agg.schema);
let agg_group_by = agg.group_by.clone();
let agg_filter_expr = agg.filter_expr.clone();
let batch_size = context.session_config().batch_size();
let input = agg.input.execute(partition, Arc::clone(&context))?;
let baseline_metrics = BaselineMetrics::new(&agg.metrics, partition);
let timer = baseline_metrics.elapsed_compute().timer();
let aggregate_exprs = agg.aggr_expr.clone();
let aggregate_arguments = aggregates::aggregate_expressions(
&agg.aggr_expr,
&agg.mode,
agg_group_by.expr.len(),
)?;
let filter_expressions = match agg.mode {
AggregateMode::Partial
| AggregateMode::Single
| AggregateMode::SinglePartitioned => agg_filter_expr,
AggregateMode::Final | AggregateMode::FinalPartitioned => {
vec![None; agg.aggr_expr.len()]
}
};
let accumulators: Vec<_> = aggregate_exprs
.iter()
.map(create_group_accumulator)
.collect::<Result<_>>()?;
let group_schema = group_schema(&agg_schema, agg_group_by.expr.len());
let name = format!("GroupedHashAggregateStream[{partition}]");
let reservation = MemoryConsumer::new(name).register(context.memory_pool());
let group_ordering = agg
.aggregation_ordering
.as_ref()
.map(|aggregation_ordering| {
GroupOrdering::try_new(&group_schema, aggregation_ordering)
})
.transpose()?
.unwrap_or(GroupOrdering::None);
let group_values = new_group_values(group_schema)?;
timer.done();
let exec_state = ExecutionState::ReadingInput;
Ok(GroupedHashAggregateStream {
schema: agg_schema,
input,
mode: agg.mode,
accumulators,
aggregate_arguments,
filter_expressions,
group_by: agg_group_by,
reservation,
group_values,
current_group_indices: Default::default(),
exec_state,
baseline_metrics,
batch_size,
group_ordering,
input_done: false,
})
}
}
fn create_group_accumulator(
agg_expr: &Arc<dyn AggregateExpr>,
) -> Result<Box<dyn GroupsAccumulator>> {
if agg_expr.groups_accumulator_supported() {
agg_expr.create_groups_accumulator()
} else {
debug!(
"Creating GroupsAccumulatorAdapter for {}: {agg_expr:?}",
agg_expr.name()
);
let agg_expr_captured = agg_expr.clone();
let factory = move || agg_expr_captured.create_accumulator();
Ok(Box::new(GroupsAccumulatorAdapter::new(factory)))
}
}
macro_rules! extract_ok {
($RES: expr) => {{
match $RES {
Ok(v) => v,
Err(e) => return Poll::Ready(Some(Err(e))),
}
}};
}
impl Stream for GroupedHashAggregateStream {
type Item = Result<RecordBatch>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
loop {
let exec_state = self.exec_state.clone();
match exec_state {
ExecutionState::ReadingInput => {
match ready!(self.input.poll_next_unpin(cx)) {
Some(Ok(batch)) => {
let timer = elapsed_compute.timer();
extract_ok!(self.group_aggregate_batch(batch));
assert!(!self.input_done);
if let Some(to_emit) = self.group_ordering.emit_to() {
let batch = extract_ok!(self.emit(to_emit));
self.exec_state = ExecutionState::ProducingOutput(batch);
}
timer.done();
}
Some(Err(e)) => {
return Poll::Ready(Some(Err(e)));
}
None => {
self.input_done = true;
self.group_ordering.input_done();
let timer = elapsed_compute.timer();
let batch = extract_ok!(self.emit(EmitTo::All));
self.exec_state = ExecutionState::ProducingOutput(batch);
timer.done();
}
}
}
ExecutionState::ProducingOutput(batch) => {
let output_batch = if batch.num_rows() <= self.batch_size {
if self.input_done {
self.exec_state = ExecutionState::Done;
} else {
self.exec_state = ExecutionState::ReadingInput
}
batch
} else {
let num_remaining = batch.num_rows() - self.batch_size;
let remaining = batch.slice(self.batch_size, num_remaining);
self.exec_state = ExecutionState::ProducingOutput(remaining);
batch.slice(0, self.batch_size)
};
return Poll::Ready(Some(Ok(
output_batch.record_output(&self.baseline_metrics)
)));
}
ExecutionState::Done => return Poll::Ready(None),
}
}
}
}
impl RecordBatchStream for GroupedHashAggregateStream {
fn schema(&self) -> SchemaRef {
self.schema.clone()
}
}
impl GroupedHashAggregateStream {
fn group_aggregate_batch(&mut self, batch: RecordBatch) -> Result<()> {
let group_by_values = evaluate_group_by(&self.group_by, &batch)?;
let input_values = evaluate_many(&self.aggregate_arguments, &batch)?;
let filter_values = evaluate_optional(&self.filter_expressions, &batch)?;
for group_values in &group_by_values {
let starting_num_groups = self.group_values.len();
self.group_values
.intern(group_values, &mut self.current_group_indices)?;
let group_indices = &self.current_group_indices;
let total_num_groups = self.group_values.len();
if total_num_groups > starting_num_groups {
self.group_ordering.new_groups(
group_values,
group_indices,
total_num_groups,
)?;
}
let t = self
.accumulators
.iter_mut()
.zip(input_values.iter())
.zip(filter_values.iter());
for ((acc, values), opt_filter) in t {
let opt_filter = opt_filter.as_ref().map(|filter| filter.as_boolean());
match self.mode {
AggregateMode::Partial
| AggregateMode::Single
| AggregateMode::SinglePartitioned => {
acc.update_batch(
values,
group_indices,
opt_filter,
total_num_groups,
)?;
}
AggregateMode::FinalPartitioned | AggregateMode::Final => {
acc.merge_batch(
values,
group_indices,
opt_filter,
total_num_groups,
)?;
}
}
}
}
self.update_memory_reservation()
}
fn update_memory_reservation(&mut self) -> Result<()> {
let acc = self.accumulators.iter().map(|x| x.size()).sum::<usize>();
self.reservation.try_resize(
acc + self.group_values.size()
+ self.group_ordering.size()
+ self.current_group_indices.allocated_size(),
)
}
fn emit(&mut self, emit_to: EmitTo) -> Result<RecordBatch> {
if self.group_values.is_empty() {
return Ok(RecordBatch::new_empty(self.schema()));
}
let mut output = self.group_values.emit(emit_to)?;
if let EmitTo::First(n) = emit_to {
self.group_ordering.remove_groups(n);
}
for acc in self.accumulators.iter_mut() {
match self.mode {
AggregateMode::Partial => output.extend(acc.state(emit_to)?),
AggregateMode::Final
| AggregateMode::FinalPartitioned
| AggregateMode::Single
| AggregateMode::SinglePartitioned => output.push(acc.evaluate(emit_to)?),
}
}
self.update_memory_reservation()?;
let batch = RecordBatch::try_new(self.schema(), output)?;
Ok(batch)
}
}