use std::collections::HashMap;
use std::marker::PhantomData;
use std::sync::Arc;
use arrow::array::{ArrayRef, BooleanArray, new_null_array};
use arrow::datatypes::SchemaRef;
use arrow::record_batch::RecordBatch;
use datafusion_common::{Result, assert_eq_or_internal_err};
use crate::aggregates::group_values::new_group_values;
use crate::aggregates::order::GroupOrdering;
use crate::aggregates::{AggregateExec, group_id_array, max_duplicate_ordinal};
use super::common::{
AggregateHashTable, AggregateHashTableBuffer, AggregateHashTableState,
EvaluatedAccumulatorArgs, HashAggregateAccumulator, PartialMarker, PartialSkipMarker,
};
impl AggregateHashTable<PartialMarker> {
pub(in crate::aggregates) fn new(
agg: &AggregateExec,
partition: usize,
output_schema: SchemaRef,
batch_size: usize,
) -> Result<Self> {
Self::new_with_filters(
agg,
partition,
Arc::clone(&output_schema),
output_schema,
batch_size,
agg.filter_expr.iter().cloned().collect(),
)
}
pub(in crate::aggregates) fn next_output_batch(
&mut self,
) -> Result<Option<RecordBatch>> {
self.next_output_batch_inner(HashAggregateAccumulator::state)
}
pub(in crate::aggregates) fn partial_skip_table(
&self,
) -> Result<AggregateHashTable<PartialSkipMarker>> {
let state = self.state.building();
let group_schema = state.group_by.group_schema(&self.input_schema)?;
let group_values = new_group_values(group_schema, &GroupOrdering::None)?;
let accumulators = state
.accumulators
.iter()
.map(HashAggregateAccumulator::empty_like)
.collect::<Result<Vec<_>>>()?;
Ok(AggregateHashTable {
group_by_metrics: self.group_by_metrics.clone(),
input_schema: Arc::clone(&self.input_schema),
output_schema: Arc::clone(&self.output_schema),
state_schema: Arc::clone(&self.state_schema),
batch_size: self.batch_size,
state: AggregateHashTableState::Building(AggregateHashTableBuffer {
group_by: Arc::clone(&state.group_by),
group_values,
batch_group_indices: Default::default(),
accumulators,
}),
_mode: PhantomData,
})
}
pub(in crate::aggregates) fn aggregate_batch(
&mut self,
batch: &RecordBatch,
) -> Result<()> {
self.aggregate_batch_inner(batch, HashAggregateAccumulator::update_batch)
}
pub(in crate::aggregates) fn start_output(&mut self) -> Result<()> {
self.init_empty_grouping_sets()?;
self.start_outputting();
Ok(())
}
fn init_empty_grouping_sets(&mut self) -> Result<()> {
let state = self.state.building_mut();
if !state.group_by.has_grouping_set() || !state.group_values.is_empty() {
return Ok(());
}
let max_ordinal = max_duplicate_ordinal(state.group_by.groups());
let mut ordinals: HashMap<&[bool], usize> = HashMap::new();
let group_schema = state.group_by.group_schema(&self.input_schema)?;
let n_expr = state.group_by.expr().len();
let mut any_interned = false;
for group in state.group_by.groups() {
let ordinal = {
let entry = ordinals.entry(group.as_slice()).or_insert(0);
let ordinal = *entry;
*entry += 1;
ordinal
};
if !group.iter().all(|&is_null| is_null) {
continue;
}
let mut cols: Vec<ArrayRef> = group_schema
.fields()
.iter()
.take(n_expr)
.map(|field| new_null_array(field.data_type(), 1))
.collect();
cols.push(group_id_array(group, ordinal, max_ordinal, 1)?);
state
.group_values
.intern(&cols, &mut state.batch_group_indices)?;
any_interned = true;
}
if any_interned {
let total_groups = state.group_values.len();
let false_filter = BooleanArray::from(vec![false]);
for acc in state.accumulators.iter_mut() {
let null_args = acc.null_arguments(&self.input_schema)?;
let values = EvaluatedAccumulatorArgs {
arguments: null_args,
filter: Some(Arc::new(false_filter.clone())),
};
acc.update_batch(&values, &[0], total_groups)?;
}
}
Ok(())
}
}
impl AggregateHashTable<PartialSkipMarker> {
pub(in crate::aggregates) fn convert_batch_to_state(
&mut self,
batch: &RecordBatch,
) -> Result<RecordBatch> {
let evaluated_batch = self.evaluate_batch(batch)?;
assert_eq_or_internal_err!(
evaluated_batch.grouping_set_args.len(),
1,
"group_values expected to have single element"
);
let mut output = evaluated_batch
.grouping_set_args
.into_iter()
.next()
.unwrap_or_default();
let state = self.state.building_mut();
for (acc, values) in state
.accumulators
.iter_mut()
.zip(evaluated_batch.accumulator_args.iter())
{
output.extend(acc.convert_to_state(values)?);
}
Ok(RecordBatch::try_new(
Arc::clone(&self.output_schema),
output,
)?)
}
}