use crate::physical_expr::EmitTo;
use arrow::row::{OwnedRow, RowConverter, Rows, SortField};
use arrow_array::ArrayRef;
use arrow_schema::Schema;
use datafusion_common::Result;
use datafusion_execution::memory_pool::proxy::VecAllocExt;
use datafusion_physical_expr::PhysicalSortExpr;
#[derive(Debug)]
pub(crate) struct GroupOrderingPartial {
state: State,
order_indices: Vec<usize>,
row_converter: RowConverter,
}
#[derive(Debug, Default)]
enum State {
#[default]
Taken,
Start,
InProgress {
current_sort: usize,
sort_key: OwnedRow,
current: usize,
},
Complete,
}
impl GroupOrderingPartial {
pub fn try_new(
input_schema: &Schema,
order_indices: &[usize],
ordering: &[PhysicalSortExpr],
) -> Result<Self> {
assert!(!order_indices.is_empty());
assert!(order_indices.len() <= ordering.len());
let fields = ordering[0..order_indices.len()]
.iter()
.map(|sort_expr| {
Ok(SortField::new_with_options(
sort_expr.expr.data_type(input_schema)?,
sort_expr.options,
))
})
.collect::<Result<Vec<_>>>()?;
Ok(Self {
state: State::Start,
order_indices: order_indices.to_vec(),
row_converter: RowConverter::new(fields)?,
})
}
fn compute_sort_keys(&mut self, group_values: &[ArrayRef]) -> Result<Rows> {
let sort_values: Vec<_> = self
.order_indices
.iter()
.map(|&idx| group_values[idx].clone())
.collect();
Ok(self.row_converter.convert_columns(&sort_values)?)
}
pub fn emit_to(&self) -> Option<EmitTo> {
match &self.state {
State::Taken => unreachable!("State previously taken"),
State::Start => None,
State::InProgress { current_sort, .. } => {
if *current_sort == 0 {
None
} else {
Some(EmitTo::First(*current_sort))
}
}
State::Complete => Some(EmitTo::All),
}
}
pub fn remove_groups(&mut self, n: usize) {
match &mut self.state {
State::Taken => unreachable!("State previously taken"),
State::Start => panic!("invalid state: start"),
State::InProgress {
current_sort,
current,
sort_key: _,
} => {
assert!(*current >= n);
*current -= n;
assert!(*current_sort >= n);
*current_sort -= n;
}
State::Complete { .. } => panic!("invalid state: complete"),
}
}
pub fn input_done(&mut self) {
self.state = match self.state {
State::Taken => unreachable!("State previously taken"),
_ => State::Complete,
};
}
pub fn new_groups(
&mut self,
batch_group_values: &[ArrayRef],
group_indices: &[usize],
total_num_groups: usize,
) -> Result<()> {
assert!(total_num_groups > 0);
assert!(!batch_group_values.is_empty());
let max_group_index = total_num_groups - 1;
let sort_keys = self.compute_sort_keys(batch_group_values)?;
let old_state = std::mem::take(&mut self.state);
let (mut current_sort, mut sort_key) = match &old_state {
State::Taken => unreachable!("State previously taken"),
State::Start => (0, sort_keys.row(0)),
State::InProgress {
current_sort,
sort_key,
..
} => (*current_sort, sort_key.row()),
State::Complete => {
panic!("Saw new group after the end of input");
}
};
let iter = group_indices.iter().zip(sort_keys.iter());
for (&group_index, group_sort_key) in iter {
if sort_key != group_sort_key {
current_sort = group_index;
sort_key = group_sort_key;
}
}
self.state = State::InProgress {
current_sort,
sort_key: sort_key.owned(),
current: max_group_index,
};
Ok(())
}
pub(crate) fn size(&self) -> usize {
std::mem::size_of::<Self>()
+ self.order_indices.allocated_size()
+ self.row_converter.size()
}
}