use std::marker::PhantomData;
use std::sync::Arc;
use arrow::array::{ArrayRef, AsArray, new_null_array};
use arrow::datatypes::SchemaRef;
use arrow::record_batch::RecordBatch;
use datafusion_common::{Result, internal_err};
use datafusion_execution::memory_pool::proxy::VecAllocExt;
use datafusion_expr::{EmitTo, GroupsAccumulator};
use datafusion_physical_expr::aggregate::AggregateFunctionExpr;
use crate::PhysicalExpr;
use crate::aggregates::group_values::{GroupByMetrics, GroupValues, new_group_values};
use crate::aggregates::grouped_hash_stream::create_group_accumulator;
use crate::aggregates::order::GroupOrdering;
use crate::aggregates::{
AggregateExec, PhysicalGroupBy, aggregate_expressions, evaluate_group_by,
};
pub(in crate::aggregates) struct PartialMarker;
pub(in crate::aggregates) struct SingleMarker;
pub(in crate::aggregates) struct PartialReduceMarker;
pub(in crate::aggregates) struct PartialSkipMarker;
pub(in crate::aggregates) struct FinalMarker;
pub(in crate::aggregates) struct AggregateHashTable<AggrMode> {
pub(super) group_by_metrics: GroupByMetrics,
pub(super) input_schema: SchemaRef,
pub(super) output_schema: SchemaRef,
pub(super) state_schema: SchemaRef,
pub(super) batch_size: usize,
pub(super) state: AggregateHashTableState,
pub(super) _mode: PhantomData<AggrMode>,
}
impl<AggrMode> AggregateHashTable<AggrMode> {
pub(super) fn new_with_filters(
agg: &AggregateExec,
partition: usize,
output_schema: SchemaRef,
state_schema: SchemaRef,
batch_size: usize,
filters: Vec<Option<Arc<dyn PhysicalExpr>>>,
) -> Result<Self> {
if batch_size == 0 {
return internal_err!("AggregateHashTable requires config batch_size >= 1");
}
let input_schema = agg.input().schema();
let aggregate_arguments = aggregate_expressions(
&agg.aggr_expr,
&agg.mode,
agg.group_by.num_group_exprs(),
)?;
let accumulators: Vec<_> = agg
.aggr_expr
.iter()
.zip(aggregate_arguments)
.zip(filters)
.map(|((agg_expr, arguments), filter)| {
let accumulator = create_group_accumulator(agg_expr)?;
Ok(HashAggregateAccumulator::new(
Arc::clone(agg_expr),
arguments,
filter,
accumulator,
))
})
.collect::<Result<_>>()?;
let group_schema = agg.group_by.group_schema(&input_schema)?;
let group_values = new_group_values(group_schema, &GroupOrdering::None)?;
Ok(Self {
group_by_metrics: GroupByMetrics::new(&agg.metrics, partition),
input_schema,
output_schema,
state_schema,
batch_size,
state: AggregateHashTableState::Building(AggregateHashTableBuffer {
group_by: Arc::clone(&agg.group_by),
group_values,
batch_group_indices: Default::default(),
accumulators,
}),
_mode: PhantomData,
})
}
pub(super) fn evaluate_batch(
&self,
batch: &RecordBatch,
) -> Result<EvaluatedAggregateBatch> {
let state = self.state.building();
let timer = self.group_by_metrics.time_calculating_group_ids.timer();
let grouping_set_args = evaluate_group_by(&state.group_by, batch)?;
drop(timer);
let timer = self.group_by_metrics.aggregate_arguments_time.timer();
let accumulator_args = self
.state
.building()
.accumulators
.iter()
.map(|acc| acc.evaluate_acc_args(batch))
.collect::<Result<Vec<_>>>()?;
drop(timer);
Ok(EvaluatedAggregateBatch {
grouping_set_args,
accumulator_args,
})
}
pub(super) fn aggregate_batch_inner(
&mut self,
batch: &RecordBatch,
aggregate_fn: AggregateBatchFn,
) -> Result<()> {
let evaluated_batch = self.evaluate_batch(batch)?;
let state = self.state.building_mut();
let _timer = self.group_by_metrics.aggregation_time.timer();
for group_values in &evaluated_batch.grouping_set_args {
state
.group_values
.intern(group_values, &mut state.batch_group_indices)?;
let group_indices = &state.batch_group_indices;
let total_num_groups = state.group_values.len();
for (acc, values) in state
.accumulators
.iter_mut()
.zip(evaluated_batch.accumulator_args.iter())
{
aggregate_fn(acc, values, group_indices, total_num_groups)?;
}
}
Ok(())
}
pub(super) fn next_output_batch_inner(
&mut self,
materialize_accumulator_fn: MaterializeAccumulatorFn,
) -> Result<Option<RecordBatch>> {
let output_schema = Arc::clone(&self.output_schema);
let batch_size = self.batch_size;
let mut output =
match std::mem::replace(&mut self.state, AggregateHashTableState::Done) {
AggregateHashTableState::Outputting(mut state) => {
if state.group_values.is_empty() {
return Ok(None);
}
let emit_to = EmitTo::All;
let timer = self.group_by_metrics.emitting_time.timer();
let mut columns = state.group_values.emit(emit_to)?;
for acc in state.accumulators.iter_mut() {
columns.extend(materialize_accumulator_fn(acc, emit_to)?);
}
drop(timer);
let batch = RecordBatch::try_new(output_schema, columns)?;
debug_assert!(batch.num_rows() > 0);
MaterializedAggregateOutput::new(batch)
}
AggregateHashTableState::OutputtingMaterialized(output) => output,
AggregateHashTableState::Done => return Ok(None),
AggregateHashTableState::Building(_) => {
return internal_err!(
"next_output_batch must be called in the outputting state"
);
}
};
let batch = output.next_batch(batch_size);
if output.is_exhausted() {
self.state = AggregateHashTableState::Done;
} else {
self.state = AggregateHashTableState::OutputtingMaterialized(output);
}
Ok(batch)
}
pub(in crate::aggregates) fn memory_size(&self) -> usize {
match &self.state {
AggregateHashTableState::Building(state)
| AggregateHashTableState::Outputting(state) => {
let acc = state
.accumulators
.iter()
.map(|acc| acc.accumulator.size())
.sum::<usize>();
acc + state.group_values.size()
+ state.batch_group_indices.allocated_size()
}
AggregateHashTableState::OutputtingMaterialized(output) => {
output.memory_size()
}
AggregateHashTableState::Done => 0,
}
}
pub(in crate::aggregates) fn group_by_metrics(&self) -> &GroupByMetrics {
&self.group_by_metrics
}
pub(in crate::aggregates) fn building_group_count(&self) -> usize {
self.state.building().group_values.len()
}
pub(in crate::aggregates) fn take_state_batch(
&mut self,
) -> Result<Option<RecordBatch>> {
let state_schema = Arc::clone(&self.state_schema);
let state = self.state.building_mut();
if state.group_values.is_empty() {
return Ok(None);
}
let mut output = state.group_values.emit(EmitTo::All)?;
for acc in &mut state.accumulators {
output.extend(acc.state(EmitTo::All)?);
}
let batch = RecordBatch::try_new(state_schema, output)?;
debug_assert!(batch.num_rows() > 0);
state.group_values.clear_shrink(0);
state.batch_group_indices.clear();
state.batch_group_indices.shrink_to_fit();
Ok(Some(batch))
}
pub(in crate::aggregates) fn is_building(&self) -> bool {
matches!(self.state, AggregateHashTableState::Building(_))
}
pub(in crate::aggregates) fn is_done(&self) -> bool {
matches!(self.state, AggregateHashTableState::Done)
}
pub(super) fn start_outputting(&mut self) {
let AggregateHashTableState::Building(mut state) =
std::mem::replace(&mut self.state, AggregateHashTableState::Done)
else {
unreachable!("hash aggregate table is not building")
};
state.batch_group_indices = Vec::new();
self.state = AggregateHashTableState::Outputting(state);
}
}
pub(super) struct HashAggregateAccumulator {
aggregate_expr: Arc<AggregateFunctionExpr>,
arguments: Vec<Arc<dyn PhysicalExpr>>,
filter: Option<Arc<dyn PhysicalExpr>>,
accumulator: Box<dyn GroupsAccumulator>,
}
pub(super) type AggregateAccumulator = HashAggregateAccumulator;
pub(super) type AggregateBatchFn = fn(
&mut AggregateAccumulator,
&EvaluatedAccumulatorArgs,
&[usize],
usize,
) -> Result<()>;
pub(super) type MaterializeAccumulatorFn =
fn(&mut AggregateAccumulator, EmitTo) -> Result<Vec<ArrayRef>>;
pub(super) struct EvaluatedAccumulatorArgs {
pub(super) arguments: Vec<ArrayRef>,
pub(super) filter: Option<ArrayRef>,
}
pub(super) struct EvaluatedAggregateBatch {
pub(super) grouping_set_args: Vec<Vec<ArrayRef>>,
pub(super) accumulator_args: Vec<EvaluatedAccumulatorArgs>,
}
pub(super) struct AggregateHashTableBuffer {
pub(super) group_by: Arc<PhysicalGroupBy>,
pub(super) group_values: Box<dyn GroupValues>,
pub(super) batch_group_indices: Vec<usize>,
pub(super) accumulators: Vec<HashAggregateAccumulator>,
}
pub(super) enum AggregateHashTableState {
Building(AggregateHashTableBuffer),
Outputting(AggregateHashTableBuffer),
OutputtingMaterialized(MaterializedAggregateOutput),
Done,
}
pub(super) struct MaterializedAggregateOutput {
batch: RecordBatch,
offset: usize,
}
impl MaterializedAggregateOutput {
pub(super) fn new(batch: RecordBatch) -> Self {
Self { batch, offset: 0 }
}
pub(super) fn next_batch(&mut self, batch_size: usize) -> Option<RecordBatch> {
debug_assert!(batch_size > 0);
if self.is_exhausted() {
return None;
}
let length = batch_size.min(self.batch.num_rows() - self.offset);
let batch = self.batch.slice(self.offset, length);
self.offset += length;
Some(batch)
}
pub(super) fn is_exhausted(&self) -> bool {
self.offset >= self.batch.num_rows()
}
pub(super) fn memory_size(&self) -> usize {
self.batch.get_array_memory_size()
}
}
impl HashAggregateAccumulator {
pub(super) fn new(
aggregate_expr: Arc<AggregateFunctionExpr>,
arguments: Vec<Arc<dyn PhysicalExpr>>,
filter: Option<Arc<dyn PhysicalExpr>>,
accumulator: Box<dyn GroupsAccumulator>,
) -> Self {
Self {
aggregate_expr,
arguments,
filter,
accumulator,
}
}
pub(super) fn empty_like(&self) -> Result<Self> {
let accumulator = create_group_accumulator(&self.aggregate_expr)?;
Ok(Self::new(
Arc::clone(&self.aggregate_expr),
self.arguments.clone(),
self.filter.clone(),
accumulator,
))
}
pub(super) fn evaluate_acc_args(
&self,
batch: &RecordBatch,
) -> Result<EvaluatedAccumulatorArgs> {
let arguments = self
.arguments
.iter()
.map(|expr| {
expr.evaluate(batch)
.and_then(|value| value.into_array(batch.num_rows()))
})
.collect::<Result<_>>()?;
let filter = self
.filter
.as_ref()
.map(|filter| {
filter
.evaluate(batch)
.and_then(|value| value.into_array(batch.num_rows()))
})
.transpose()?;
Ok(EvaluatedAccumulatorArgs { arguments, filter })
}
pub(super) fn size(&self) -> usize {
self.accumulator.size()
}
pub(super) fn update_batch(
&mut self,
values: &EvaluatedAccumulatorArgs,
group_indices: &[usize],
total_num_groups: usize,
) -> Result<()> {
let filter = values.filter.as_ref().map(|filter| filter.as_boolean());
self.accumulator.update_batch(
&values.arguments,
group_indices,
filter,
total_num_groups,
)
}
pub(super) fn merge_batch(
&mut self,
values: &EvaluatedAccumulatorArgs,
group_indices: &[usize],
total_num_groups: usize,
) -> Result<()> {
debug_assert!(values.filter.is_none());
self.accumulator
.merge_batch(&values.arguments, group_indices, total_num_groups)
}
pub(super) fn evaluate(&mut self, emit_to: EmitTo) -> Result<ArrayRef> {
self.accumulator.evaluate(emit_to)
}
pub(super) fn evaluate_to_columns(
&mut self,
emit_to: EmitTo,
) -> Result<Vec<ArrayRef>> {
Ok(vec![self.evaluate(emit_to)?])
}
pub(super) fn state(&mut self, emit_to: EmitTo) -> Result<Vec<ArrayRef>> {
self.accumulator.state(emit_to)
}
pub(super) fn convert_to_state(
&mut self,
values: &EvaluatedAccumulatorArgs,
) -> Result<Vec<ArrayRef>> {
let opt_filter = values.filter.as_ref().map(|filter| filter.as_boolean());
self.accumulator
.convert_to_state(&values.arguments, opt_filter)
}
pub(super) fn null_arguments(
&self,
input_schema: &SchemaRef,
) -> Result<Vec<ArrayRef>> {
self.arguments
.iter()
.map(|expr| {
let data_type = expr.data_type(input_schema)?;
Ok(new_null_array(&data_type, 1))
})
.collect()
}
}
impl AggregateHashTableState {
pub(super) fn building(&self) -> &AggregateHashTableBuffer {
let Self::Building(state) = self else {
unreachable!("hash aggregate table is not building")
};
state
}
pub(super) fn building_mut(&mut self) -> &mut AggregateHashTableBuffer {
let Self::Building(state) = self else {
unreachable!("hash aggregate table is not building")
};
state
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use arrow::array::{Array, Int32Array};
use arrow::datatypes::{DataType, Field, Schema};
use super::*;
#[test]
fn materialized_aggregate_output_slices_batches_until_exhausted() -> Result<()> {
let schema = Arc::new(Schema::new(vec![Field::new(
"group_col",
DataType::Int32,
false,
)]));
let batch = RecordBatch::try_new(
schema,
vec![Arc::new(Int32Array::from(vec![1, 2, 3, 4, 5]))],
)?;
let mut output = MaterializedAggregateOutput::new(batch);
assert_eq!(int32_values(&output.next_batch(2).unwrap(), 0), vec![1, 2]);
assert_eq!(int32_values(&output.next_batch(2).unwrap(), 0), vec![3, 4]);
assert_eq!(int32_values(&output.next_batch(2).unwrap(), 0), vec![5]);
assert!(output.next_batch(2).is_none());
assert!(output.is_exhausted());
Ok(())
}
fn int32_values(batch: &RecordBatch, column: usize) -> Vec<i32> {
let array = batch
.column(column)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
(0..array.len()).map(|idx| array.value(idx)).collect()
}
}