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::{DataFusionError, Result, internal_datafusion_err, internal_err};
use datafusion_execution::TaskContext;
use datafusion_execution::memory_pool::{MemoryConsumer, MemoryReservation};
use datafusion_physical_expr::PhysicalSortExpr;
use datafusion_physical_expr::expressions::Column;
use datafusion_physical_expr_common::sort_expr::LexOrdering;
use futures::stream::{Stream, StreamExt};
use super::aggregate_hash_table::{AggregateHashTable, SingleMarker};
use super::group_values::GroupByMetrics;
use super::ordered_final_stream::OrderedFinalAggregateStream;
use super::{AggregateExec, create_schema};
use crate::aggregates::AggregateMode;
use crate::metrics::{BaselineMetrics, RecordOutput, SpillMetrics};
use crate::sorts::IncrementalSortIterator;
use crate::sorts::streaming_merge::{SortedSpillFile, StreamingMergeBuilder};
use crate::spill::spill_manager::SpillManager;
use crate::stream::EmptyRecordBatchStream;
use crate::{InputOrderMode, RecordBatchStream, SendableRecordBatchStream};
pub(crate) struct SingleHashAggregateStream {
schema: SchemaRef,
input: SendableRecordBatchStream,
baseline_metrics: BaselineMetrics,
reservation: MemoryReservation,
state: Option<SingleHashAggregateState>,
}
struct SingleSpillContext {
final_agg: AggregateExec,
context: Arc<TaskContext>,
partition: usize,
batch_size: usize,
spill_expr: LexOrdering,
spill_manager: SpillManager,
spills: Vec<SortedSpillFile>,
}
enum SingleHashAggregateState {
ReadingInput {
hash_table: AggregateHashTable<SingleMarker>,
spill_context: Option<Box<SingleSpillContext>>,
},
Spilling {
hash_table: AggregateHashTable<SingleMarker>,
spill_context: Box<SingleSpillContext>,
},
ProducingOutput {
hash_table: AggregateHashTable<SingleMarker>,
},
PreparingMergeInput {
hash_table: AggregateHashTable<SingleMarker>,
spill_context: Box<SingleSpillContext>,
},
MergingSpills {
stream: SendableRecordBatchStream,
},
Done,
Error,
}
type SingleHashAggregatePoll = Poll<Option<Result<RecordBatch>>>;
type SingleHashAggregateStateTransition = ControlFlow<
(SingleHashAggregatePoll, SingleHashAggregateState),
SingleHashAggregateState,
>;
impl SingleSpillContext {
fn new(
agg: &AggregateExec,
context: &Arc<TaskContext>,
partition: usize,
batch_size: usize,
spill_schema: &SchemaRef,
spill_metrics: SpillMetrics,
) -> Result<Self> {
let group_schema = agg.group_by.group_schema(&agg.input().schema())?;
let output_ordering = agg.cache.output_ordering();
let spill_sort_exprs =
group_schema
.fields()
.iter()
.enumerate()
.map(|(idx, field)| {
let output_expr = Column::new(field.name(), idx);
let sort_options = output_ordering
.and_then(|ordering| ordering.get_sort_options(&output_expr))
.unwrap_or_default();
PhysicalSortExpr::new(Arc::new(output_expr), sort_options)
});
let Some(spill_expr) = LexOrdering::new(spill_sort_exprs) else {
return internal_err!("Single hash aggregate spill expression is empty");
};
let spill_manager = SpillManager::new(
context.runtime_env(),
spill_metrics,
Arc::clone(spill_schema),
)
.with_compression_type(context.session_config().spill_compression());
let mut final_agg = agg.clone();
final_agg.mode = match agg.mode {
AggregateMode::Single => AggregateMode::Final,
AggregateMode::SinglePartitioned => AggregateMode::FinalPartitioned,
mode => {
return internal_err!(
"Single hash aggregate spill cannot replay aggregate mode {mode:?}"
);
}
};
final_agg.group_by = Arc::new(agg.group_by.as_final());
final_agg.input_order_mode = InputOrderMode::Sorted;
Ok(Self {
final_agg,
context: Arc::clone(context),
partition,
batch_size,
spill_expr,
spill_manager,
spills: vec![],
})
}
fn has_spills(&self) -> bool {
!self.spills.is_empty()
}
fn spill_table(
&mut self,
hash_table: &mut AggregateHashTable<SingleMarker>,
) -> Result<()> {
let Some(batch) = hash_table.take_state_batch()? else {
return Ok(());
};
let sorted_iter =
IncrementalSortIterator::new(batch, self.spill_expr.clone(), self.batch_size);
let spill_file = self
.spill_manager
.spill_record_batch_iter_and_return_max_batch_memory(
sorted_iter,
"SingleHashAggregateSpill",
)?;
let Some((file, max_record_batch_memory)) = spill_file else {
return internal_err!("Single hash aggregation produced an empty spill");
};
self.spills.push(SortedSpillFile {
file,
max_record_batch_memory,
});
Ok(())
}
fn into_replay_stream(
self,
baseline_metrics: &BaselineMetrics,
group_by_metrics: GroupByMetrics,
reservation: MemoryReservation,
) -> Result<SendableRecordBatchStream> {
let Self {
final_agg,
context,
partition,
batch_size,
spill_expr,
spill_manager,
spills,
} = self;
let spill_schema = Arc::clone(spill_manager.schema());
let merge_reservation = reservation.new_empty();
let merged = StreamingMergeBuilder::new()
.with_schema(spill_schema)
.with_spill_manager(spill_manager)
.with_sorted_spill_files(spills)
.with_expressions(&spill_expr)
.with_metrics(baseline_metrics.intermediate())
.with_batch_size(batch_size)
.with_reservation(merge_reservation)
.build()?;
let replay = OrderedFinalAggregateStream::new_with_input_and_metrics(
&final_agg,
&context,
partition,
merged,
&InputOrderMode::Sorted,
baseline_metrics.clone(),
group_by_metrics,
None,
reservation,
)?;
Ok(Box::pin(replay))
}
}
impl SingleHashAggregateStream {
pub fn new(
agg: &AggregateExec,
context: &Arc<TaskContext>,
partition: usize,
) -> Result<Self> {
debug_assert!(matches!(
agg.mode,
AggregateMode::Single | AggregateMode::SinglePartitioned
));
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 input_schema = input.schema();
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 state_schema = Arc::new(create_schema(
input_schema.as_ref(),
&agg.group_by,
&agg.aggr_expr,
AggregateMode::Partial,
)?);
let hash_table = AggregateHashTable::<SingleMarker>::new(
agg,
partition,
Arc::clone(&schema),
Arc::clone(&state_schema),
batch_size,
)?;
let can_spill = context.runtime_env().disk_manager.tmp_files_enabled();
let spill_context = if can_spill {
Some(Box::new(SingleSpillContext::new(
agg,
context,
partition,
batch_size,
&state_schema,
spill_metrics,
)?))
} else {
None
};
let reservation =
MemoryConsumer::new(format!("SingleHashAggregateStream[{partition}]"))
.with_can_spill(can_spill)
.register(context.memory_pool());
Ok(Self {
schema,
input,
baseline_metrics,
reservation,
state: Some(SingleHashAggregateState::ReadingInput {
hash_table,
spill_context,
}),
})
}
fn close_input(&mut self) {
let input_schema = self.input.schema();
self.input = Box::pin(EmptyRecordBatchStream::new(input_schema));
}
fn break_with_err(error: DataFusionError) -> SingleHashAggregateStateTransition {
ControlFlow::Break((
Poll::Ready(Some(Err(error))),
SingleHashAggregateState::Error,
))
}
fn break_with_internal_err(message: &str) -> SingleHashAggregateStateTransition {
Self::break_with_err(internal_datafusion_err!("{message}"))
}
fn reservation_size_for_table(
hash_table: &AggregateHashTable<SingleMarker>,
spill_context: Option<&SingleSpillContext>,
) -> usize {
let table_size = hash_table.memory_size();
if spill_context.is_some() {
table_size.saturating_add(
hash_table
.building_group_count()
.saturating_mul(size_of::<u32>()),
)
} else {
table_size
}
}
fn handle_reading_input(
&mut self,
cx: &mut Context<'_>,
original_state: SingleHashAggregateState,
) -> SingleHashAggregateStateTransition {
let SingleHashAggregateState::ReadingInput {
mut hash_table,
spill_context,
} = original_state
else {
return Self::break_with_internal_err(
"Single hash aggregate stream expected ReadingInput state",
);
};
match self.input.poll_next_unpin(cx) {
Poll::Pending => ControlFlow::Break((
Poll::Pending,
SingleHashAggregateState::ReadingInput {
hash_table,
spill_context,
},
)),
Poll::Ready(Some(Ok(batch))) => {
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
let timer = elapsed_compute.timer();
let result = hash_table.aggregate_batch(&batch);
timer.done();
if let Err(e) = result {
return Self::break_with_err(e);
}
let timer = elapsed_compute.timer();
let resize_result =
self.reservation
.try_resize(Self::reservation_size_for_table(
&hash_table,
spill_context.as_deref(),
));
timer.done();
match resize_result {
Ok(()) => {}
Err(e @ DataFusionError::ResourcesExhausted(_)) => {
let Some(spill_context) = spill_context else {
return Self::break_with_err(e.context(
"Single hash aggregate cannot spill because temporary files are not enabled in the DiskManager",
));
};
if hash_table.building_group_count() == 0 {
return Self::break_with_internal_err(
"Single hash aggregate ran out of memory with no aggregated groups",
);
}
return ControlFlow::Continue(
SingleHashAggregateState::Spilling {
hash_table,
spill_context,
},
);
}
Err(e) => {
return Self::break_with_err(e);
}
}
ControlFlow::Continue(SingleHashAggregateState::ReadingInput {
hash_table,
spill_context,
})
}
Poll::Ready(Some(Err(e))) => Self::break_with_err(e),
Poll::Ready(None) => {
self.close_input();
match spill_context {
Some(spill_context) if spill_context.has_spills() => {
ControlFlow::Continue(
SingleHashAggregateState::PreparingMergeInput {
hash_table,
spill_context,
},
)
}
_ => {
let elapsed_compute =
self.baseline_metrics.elapsed_compute().clone();
let timer = elapsed_compute.timer();
let result = hash_table.start_output();
timer.done();
match result {
Ok(()) => ControlFlow::Continue(
SingleHashAggregateState::ProducingOutput { hash_table },
),
Err(e) => Self::break_with_err(e),
}
}
}
}
}
}
fn handle_spilling(
&mut self,
original_state: SingleHashAggregateState,
) -> SingleHashAggregateStateTransition {
let SingleHashAggregateState::Spilling {
mut hash_table,
mut spill_context,
} = original_state
else {
return Self::break_with_internal_err(
"Single hash aggregate stream expected Spilling state",
);
};
if hash_table.building_group_count() == 0 {
return Self::break_with_internal_err(
"Single hash aggregation entered Spilling with an empty table",
);
}
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
let timer = elapsed_compute.timer();
let mut result = spill_context.spill_table(&mut hash_table);
if let Err(e) = self.reservation.try_resize(hash_table.memory_size()) {
result =
Err(e.context("Decreasing allocation after spilling should succeed"));
}
timer.done();
match result {
Ok(()) => ControlFlow::Continue(SingleHashAggregateState::ReadingInput {
hash_table,
spill_context: Some(spill_context),
}),
Err(e) => Self::break_with_err(e),
}
}
fn handle_preparing_merge_input(
&mut self,
original_state: SingleHashAggregateState,
) -> SingleHashAggregateStateTransition {
let SingleHashAggregateState::PreparingMergeInput {
mut hash_table,
mut spill_context,
} = original_state
else {
return Self::break_with_internal_err(
"Single hash aggregate stream expected PreparingMergeInput state",
);
};
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
let timer = elapsed_compute.timer();
let replay = match spill_context.spill_table(&mut hash_table) {
Ok(()) => {
let group_by_metrics = hash_table.group_by_metrics().clone();
drop(hash_table);
match self.reservation.try_resize(0) {
Ok(()) => (*spill_context).into_replay_stream(
&self.baseline_metrics,
group_by_metrics,
self.reservation.new_empty(),
),
Err(e) => Err(e),
}
}
Err(e) => Err(e),
};
timer.done();
match replay {
Ok(stream) => {
ControlFlow::Continue(SingleHashAggregateState::MergingSpills { stream })
}
Err(e) => Self::break_with_err(e),
}
}
fn handle_merging_spills(
&mut self,
cx: &mut Context<'_>,
original_state: SingleHashAggregateState,
) -> SingleHashAggregateStateTransition {
let SingleHashAggregateState::MergingSpills { mut stream } = original_state
else {
return Self::break_with_internal_err(
"Single hash aggregate stream expected MergingSpills state",
);
};
match stream.poll_next_unpin(cx) {
Poll::Pending => ControlFlow::Break((
Poll::Pending,
SingleHashAggregateState::MergingSpills { stream },
)),
Poll::Ready(Some(Ok(batch))) => ControlFlow::Break((
Poll::Ready(Some(Ok(batch))),
SingleHashAggregateState::MergingSpills { stream },
)),
Poll::Ready(Some(Err(e))) => Self::break_with_err(e),
Poll::Ready(None) => ControlFlow::Continue(SingleHashAggregateState::Done),
}
}
fn handle_producing_output(
&mut self,
original_state: SingleHashAggregateState,
) -> SingleHashAggregateStateTransition {
let SingleHashAggregateState::ProducingOutput { mut hash_table } = original_state
else {
return Self::break_with_internal_err(
"Single hash aggregate stream expected ProducingOutput state",
);
};
let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
let timer = elapsed_compute.timer();
let result = hash_table.next_output_batch();
timer.done();
match result {
Ok(Some(batch)) => {
let next_state = if hash_table.is_done() {
drop(hash_table);
if let Err(e) = self.reservation.try_resize(0) {
return Self::break_with_err(e);
}
SingleHashAggregateState::Done
} else {
if let Err(e) = self.reservation.try_resize(hash_table.memory_size())
{
return Self::break_with_err(e);
}
SingleHashAggregateState::ProducingOutput { hash_table }
};
ControlFlow::Break((
Poll::Ready(Some(Ok(batch.record_output(&self.baseline_metrics)))),
next_state,
))
}
Err(e) => Self::break_with_err(e),
Ok(None) => {
drop(hash_table);
let next_state = SingleHashAggregateState::Done;
if let Err(e) = self.reservation.try_resize(0) {
return Self::break_with_err(e);
}
ControlFlow::Continue(next_state)
}
}
}
}
impl Stream for SingleHashAggregateStream {
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("SingleHashAggregateStream state should not be None");
let next_state = match cur_state {
state @ SingleHashAggregateState::ReadingInput { .. } => {
self.handle_reading_input(cx, state)
}
state @ SingleHashAggregateState::Spilling { .. } => {
self.handle_spilling(state)
}
state @ SingleHashAggregateState::PreparingMergeInput { .. } => {
self.handle_preparing_merge_input(state)
}
state @ SingleHashAggregateState::MergingSpills { .. } => {
self.handle_merging_spills(cx, state)
}
state @ SingleHashAggregateState::ProducingOutput { .. } => {
self.handle_producing_output(state)
}
state @ SingleHashAggregateState::Error => {
self.close_input();
self.reservation.free();
self.state = Some(state);
return Poll::Ready(None);
}
state @ SingleHashAggregateState::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::Ready(Some(Err(e))), next_state)) => {
debug_assert!(matches!(next_state, SingleHashAggregateState::Error));
self.close_input();
self.reservation.free();
self.state = Some(SingleHashAggregateState::Error);
return Poll::Ready(Some(Err(e)));
}
ControlFlow::Break((poll, next_state)) => {
self.state = Some(next_state);
return poll;
}
}
}
}
}
impl RecordBatchStream for SingleHashAggregateStream {
fn schema(&self) -> SchemaRef {
Arc::clone(&self.schema)
}
}