use std::sync::Arc;
use arrow::datatypes::SchemaRef;
use arrow::record_batch::RecordBatch;
use datafusion_common::{DataFusionError, Result};
use datafusion_execution::memory_pool::{MemoryConsumer, MemoryReservation};
use datafusion_execution::{TaskContext, TryEmitter, async_try_stream};
use futures::stream::{Stream, StreamExt};
use super::AggregateExec;
use super::aggregate_hash_table::{OrderedAggregateTable, PartialMarker};
use crate::aggregates::AggregateMode;
use crate::aggregates::order::GroupOrdering;
use crate::metrics::{BaselineMetrics, MetricBuilder, SpillMetrics};
use crate::stream::{EmptyRecordBatchStream, ObservedStream, RecordBatchStreamAdapter};
use crate::{InputOrderMode, SendableRecordBatchStream, metrics};
pub(crate) struct OrderedPartialAggregateStream {
schema: SchemaRef,
input: SendableRecordBatchStream,
reservation: MemoryReservation,
baseline_metrics: BaselineMetrics,
reduction_factor: metrics::RatioMetrics,
table: Option<OrderedAggregateTable<PartialMarker>>,
}
impl OrderedPartialAggregateStream {
pub fn new(
agg: &AggregateExec,
context: &Arc<TaskContext>,
partition: usize,
) -> Result<Self> {
debug_assert_eq!(agg.mode, AggregateMode::Partial);
debug_assert_ne!(agg.input_order_mode, InputOrderMode::Linear);
let schema = Arc::clone(&agg.schema);
let input = agg.input.execute(partition, Arc::clone(context))?;
let batch_size = context.session_config().batch_size();
let baseline_metrics = BaselineMetrics::new(&agg.metrics, partition);
let _spill_metrics = SpillMetrics::new(&agg.metrics, partition);
let reduction_factor = MetricBuilder::new(&agg.metrics)
.with_type(metrics::MetricType::Summary)
.ratio_metrics("reduction_factor", partition);
let table = OrderedAggregateTable::<PartialMarker>::new(
agg,
partition,
Arc::clone(&schema),
batch_size,
)?;
let reservation =
MemoryConsumer::new(format!("OrderedPartialAggregateStream[{partition}]"))
.with_can_spill(matches!(
table.group_ordering(),
GroupOrdering::Partial(_)
))
.register(context.memory_pool());
Ok(Self {
schema,
input,
reservation,
baseline_metrics,
reduction_factor,
table: Some(table),
})
}
pub(crate) fn into_stream(self) -> SendableRecordBatchStream {
let schema_clone = Arc::clone(&self.schema);
let cloned_metrics = self.baseline_metrics.clone();
let stream = Box::pin(RecordBatchStreamAdapter::new(
schema_clone,
self.create_stream(),
));
Box::pin(ObservedStream::new(stream, cloned_metrics, None))
}
fn create_stream(mut self) -> impl Stream<Item = Result<RecordBatch>> {
async_try_stream(|mut emitter| async move {
let mut table = self
.table
.take()
.expect("OrderedPartialAggregateStream state should not be None");
self.handle_reading_input(&mut table, &mut emitter).await?;
self.close_input();
table.input_done();
self.handle_draining_final(table, &mut emitter).await?;
Ok(())
})
}
fn close_input(&mut self) {
let input_schema = self.input.schema();
self.input = Box::pin(EmptyRecordBatchStream::new(input_schema));
}
async fn handle_reading_input(
&mut self,
table: &mut OrderedAggregateTable<PartialMarker>,
emitter: &mut TryEmitter<RecordBatch, DataFusionError>,
) -> Result<()> {
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
while let Some(batch) = self.input.next().await.transpose()? {
let input_rows = batch.num_rows();
self.reduction_factor.add_total(input_rows);
let timer = elapsed_compute.timer();
table.aggregate_batch(&batch)?;
if let Some(batch) = self.resize_or_take_state_batch(table)? {
self.reduction_factor.add_part(batch.num_rows());
drop(timer);
emitter.emit(batch).await;
continue;
}
let Some(batch) = table.next_output_batch()? else {
continue;
};
self.reduction_factor.add_part(batch.num_rows());
self.reservation.try_resize(table.memory_size())?;
drop(timer);
emitter.emit(batch).await;
}
Ok(())
}
fn resize_or_take_state_batch(
&mut self,
table: &mut OrderedAggregateTable<PartialMarker>,
) -> Result<Option<RecordBatch>> {
let oom = match self.reservation.try_resize(table.memory_size()) {
Ok(()) => return Ok(None),
Err(e @ DataFusionError::ResourcesExhausted(_)) => e,
Err(e) => return Err(e),
};
if matches!(table.group_ordering(), GroupOrdering::Full(_)) {
return Err(oom);
}
let Some(batch) = table.take_state_batch()? else {
return Err(oom);
};
self.reservation.try_resize(table.memory_size())?;
Ok(Some(batch))
}
async fn handle_draining_final(
&mut self,
mut table: OrderedAggregateTable<PartialMarker>,
emitter: &mut TryEmitter<RecordBatch, DataFusionError>,
) -> Result<()> {
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
let mut timer = elapsed_compute.timer();
while let Some(batch) = table.next_output_batch()? {
self.reduction_factor.add_part(batch.num_rows());
if table.is_empty() {
drop(table);
let _ = self.reservation.try_resize(0);
drop(timer);
emitter.emit(batch).await;
return Ok(());
}
self.reservation.try_resize(table.memory_size())?;
timer.done();
emitter.emit(batch).await;
timer = elapsed_compute.timer();
}
Ok(())
}
}