use std::ops::ControlFlow;
use std::sync::Arc;
use std::task::{Context, Poll};
use arrow::datatypes::SchemaRef;
use arrow::record_batch::RecordBatch;
use datafusion_common::Result;
use datafusion_execution::TaskContext;
use datafusion_execution::memory_pool::{MemoryConsumer, MemoryReservation};
use futures::stream::{Stream, StreamExt};
use super::AggregateExec;
use super::aggregate_hash_table::{AggregateHashTable, PartialReduceMarker};
use crate::metrics::{BaselineMetrics, RecordOutput, SpillMetrics};
use crate::stream::EmptyRecordBatchStream;
use crate::{InputOrderMode, RecordBatchStream, SendableRecordBatchStream};
pub(crate) struct PartialReduceHashAggregateStream {
schema: SchemaRef,
input: SendableRecordBatchStream,
baseline_metrics: BaselineMetrics,
reservation: MemoryReservation,
state: Option<PartialReduceHashAggregateState>,
}
enum PartialReduceHashAggregateState {
ReadingInput {
hash_table: AggregateHashTable<PartialReduceMarker>,
},
ProducingOutput {
hash_table: AggregateHashTable<PartialReduceMarker>,
},
Done,
}
type PartialReduceHashAggregatePoll = Poll<Option<Result<RecordBatch>>>;
type PartialReduceHashAggregateStateTransition = ControlFlow<
(
PartialReduceHashAggregatePoll,
PartialReduceHashAggregateState,
),
PartialReduceHashAggregateState,
>;
impl PartialReduceHashAggregateState {
fn hash_table(&self) -> &AggregateHashTable<PartialReduceMarker> {
match self {
Self::ReadingInput { hash_table } | Self::ProducingOutput { hash_table } => {
hash_table
}
Self::Done => unreachable!("Done state does not hold a hash table"),
}
}
fn hash_table_mut(&mut self) -> &mut AggregateHashTable<PartialReduceMarker> {
match self {
Self::ReadingInput { hash_table } | Self::ProducingOutput { hash_table } => {
hash_table
}
Self::Done => unreachable!("Done state does not hold a hash table"),
}
}
fn into_hash_table(self) -> AggregateHashTable<PartialReduceMarker> {
match self {
Self::ReadingInput { hash_table } | Self::ProducingOutput { hash_table } => {
hash_table
}
Self::Done => unreachable!("Done state does not hold a hash table"),
}
}
fn into_producing_output(self) -> Self {
Self::ProducingOutput {
hash_table: self.into_hash_table(),
}
}
fn into_done(self) -> Self {
Self::Done
}
}
impl PartialReduceHashAggregateStream {
pub fn new(
agg: &AggregateExec,
context: &Arc<TaskContext>,
partition: usize,
) -> Result<Self> {
debug_assert_eq!(agg.mode, super::AggregateMode::PartialReduce);
debug_assert_eq!(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 hash_table = AggregateHashTable::<PartialReduceMarker>::new(
agg,
partition,
Arc::clone(&schema),
batch_size,
)?;
let reservation =
MemoryConsumer::new(format!("PartialReduceHashAggregateStream[{partition}]"))
.register(context.memory_pool());
Ok(Self {
schema,
input,
baseline_metrics,
reservation,
state: Some(PartialReduceHashAggregateState::ReadingInput { hash_table }),
})
}
fn start_output(
&mut self,
hash_table: &mut AggregateHashTable<PartialReduceMarker>,
) -> Result<()> {
let input_schema = self.input.schema();
self.input = Box::pin(EmptyRecordBatchStream::new(input_schema));
hash_table.start_output()
}
fn handle_reading_input(
&mut self,
cx: &mut Context<'_>,
mut original_state: PartialReduceHashAggregateState,
) -> PartialReduceHashAggregateStateTransition {
debug_assert!(matches!(
&original_state,
PartialReduceHashAggregateState::ReadingInput { .. }
));
debug_assert!(original_state.hash_table().is_building());
match self.input.poll_next_unpin(cx) {
Poll::Pending => ControlFlow::Break((Poll::Pending, original_state)),
Poll::Ready(Some(Ok(batch))) => {
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
let timer = elapsed_compute.timer();
let result = original_state.hash_table_mut().aggregate_batch(&batch);
timer.done();
if let Err(e) = result {
return ControlFlow::Break((
Poll::Ready(Some(Err(e))),
original_state,
));
}
if let Err(e) = self
.reservation
.try_resize(original_state.hash_table().memory_size())
{
return ControlFlow::Break((
Poll::Ready(Some(Err(e))),
original_state,
));
}
ControlFlow::Continue(original_state)
}
Poll::Ready(Some(Err(e))) => {
ControlFlow::Break((Poll::Ready(Some(Err(e))), original_state))
}
Poll::Ready(None) => {
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
let timer = elapsed_compute.timer();
let result = self.start_output(original_state.hash_table_mut());
timer.done();
match result {
Ok(()) => {
ControlFlow::Continue(original_state.into_producing_output())
}
Err(e) => {
ControlFlow::Break((Poll::Ready(Some(Err(e))), original_state))
}
}
}
}
}
fn handle_producing_output(
&mut self,
mut original_state: PartialReduceHashAggregateState,
) -> PartialReduceHashAggregateStateTransition {
debug_assert!(matches!(
&original_state,
PartialReduceHashAggregateState::ProducingOutput { .. }
));
debug_assert!(!original_state.hash_table().is_building());
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
let timer = elapsed_compute.timer();
let result = original_state.hash_table_mut().next_output_batch();
timer.done();
match result {
Ok(Some(batch)) => {
let _ = self
.reservation
.try_resize(original_state.hash_table().memory_size());
debug_assert!(batch.num_rows() > 0);
let next_state = if original_state.hash_table().is_done() {
original_state.into_done()
} else {
original_state
};
ControlFlow::Break((
Poll::Ready(Some(Ok(batch.record_output(&self.baseline_metrics)))),
next_state,
))
}
Ok(None) => {
let _ = self.reservation.try_resize(0);
ControlFlow::Continue(original_state.into_done())
}
Err(e) => ControlFlow::Break((Poll::Ready(Some(Err(e))), original_state)),
}
}
}
impl Stream for PartialReduceHashAggregateStream {
type Item = Result<RecordBatch>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
loop {
let cur_state = self
.state
.take()
.expect("PartialReduceHashAggregateStream state should not be None");
let next_state = match cur_state {
state @ PartialReduceHashAggregateState::ReadingInput { .. } => {
self.handle_reading_input(cx, state)
}
state @ PartialReduceHashAggregateState::ProducingOutput { .. } => {
self.handle_producing_output(state)
}
state @ PartialReduceHashAggregateState::Done => {
let _ = self.reservation.try_resize(0);
self.state = Some(state);
return Poll::Ready(None);
}
};
match next_state {
ControlFlow::Continue(next_state) => {
self.state = Some(next_state);
continue;
}
ControlFlow::Break((poll, next_state)) => {
self.state = Some(next_state);
return poll;
}
}
}
}
}
impl RecordBatchStream for PartialReduceHashAggregateStream {
fn schema(&self) -> SchemaRef {
Arc::clone(&self.schema)
}
}