use std::borrow::{Borrow, BorrowMut};
use glaredb_error::Result;
use super::aggregate_layout::AggregateLayout;
use super::block::ValidityInitializer;
use super::row_blocks::{BlockAppendState, RowBlocks};
use super::row_scan::RowScanState;
use crate::arrays::array::Array;
use crate::buffer::buffer_manager::DefaultBufferManager;
use crate::util::iter::IntoExactSizeIterator;
#[derive(Debug)]
pub struct AggregateAppendState {
block_append: BlockAppendState,
heap_sizes: Vec<usize>,
}
impl AggregateAppendState {
pub fn row_pointers(&self) -> &[*mut u8] {
&self.block_append.row_pointers
}
}
#[derive(Debug)]
pub struct AggregateCollection {
layout: AggregateLayout,
blocks: RowBlocks<ValidityInitializer>,
}
impl AggregateCollection {
pub fn new(layout: AggregateLayout, block_capacity: usize) -> Self {
let blocks = RowBlocks::new(
&DefaultBufferManager,
ValidityInitializer::from_aggregate_layout(&layout),
layout.row_width,
block_capacity,
Some(layout.base_align),
);
AggregateCollection { layout, blocks }
}
pub fn init_append_state(&self) -> AggregateAppendState {
AggregateAppendState {
block_append: BlockAppendState {
row_pointers: Vec::new(),
heap_pointers: Vec::new(),
},
heap_sizes: Vec::new(),
}
}
pub fn row_mut_ptr_iter(&self) -> impl Iterator<Item = *mut u8> + '_ {
self.blocks.row_mut_ptr_iter()
}
pub fn num_groups(&self) -> usize {
self.blocks.total_rows()
}
pub fn num_row_blocks(&self) -> usize {
self.blocks.num_row_blocks()
}
pub(crate) fn row_blocks(&self) -> &RowBlocks<ValidityInitializer> {
&self.blocks
}
pub(crate) fn append_groups<A>(
&mut self,
state: &mut AggregateAppendState,
groups: &[A],
rows: impl IntoExactSizeIterator<Item = usize> + Clone,
) -> Result<()>
where
A: Borrow<Array>,
{
debug_assert_eq!(groups.len(), self.layout.groups.num_columns());
let num_rows = rows.clone().into_exact_size_iter().len();
state.block_append.clear();
if self.layout.groups.requires_heap {
state.heap_sizes.resize(num_rows, 0);
self.layout
.groups
.compute_heap_sizes(groups, rows.clone(), &mut state.heap_sizes)?;
}
if self.layout.groups.requires_heap {
self.blocks.prepare_append(
&mut state.block_append,
num_rows,
Some(&state.heap_sizes),
)?
} else {
self.blocks
.prepare_append(&mut state.block_append, num_rows, None)?
}
unsafe {
self.layout
.groups
.write_arrays(&mut state.block_append, groups, rows)?
};
for (offset, agg) in self.layout.iter_offsets_and_aggregates() {
for row_ptr in &state.block_append.row_pointers {
unsafe {
let state_ptr = row_ptr.byte_add(offset);
agg.function.call_new_aggregate_state(state_ptr)
}
}
}
Ok(())
}
pub fn scan_groups<A>(
&self,
state: &mut RowScanState,
outputs: &mut [A],
count: usize,
) -> Result<usize>
where
A: BorrowMut<Array>,
{
state.scan(&self.layout.groups, &self.blocks, outputs, count)
}
pub fn scan_groups_subset<A>(
&self,
state: &mut RowScanState,
columns: impl IntoExactSizeIterator<Item = usize> + Clone,
outputs: &mut [A],
count: usize,
) -> Result<usize>
where
A: BorrowMut<Array>,
{
state.scan_subset(&self.layout.groups, &self.blocks, columns, outputs, count)
}
#[allow(unused)] pub(crate) unsafe fn finalize_groups<A>(
&self,
group_ptrs: &mut [*mut u8],
groups: &mut [A],
results: &mut [A],
) -> Result<()>
where
A: BorrowMut<Array>,
{
unsafe {
debug_assert_eq!(groups.len(), self.layout.groups.num_columns());
debug_assert_eq!(results.len(), self.layout.aggregates.len());
{
let group_ptrs = group_ptrs.iter().copied().map(|ptr| ptr as _);
self.layout
.groups
.read_arrays(group_ptrs, groups.iter_mut().enumerate(), 0)?;
}
self.layout.finalize_states(group_ptrs, results)?;
Ok(())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::arrays::datatype::DataType;
use crate::arrays::row::aggregate_layout::AggregateUpdateSelector;
use crate::expr::physical::PhysicalAggregateExpression;
use crate::expr::{self, bind_aggregate_function};
use crate::functions::aggregate::builtin::sum::FUNCTION_SET_SUM;
use crate::testutil::arrays::assert_arrays_eq;
use crate::util::iter::TryFromExactSizeIterator;
#[test]
fn append_groups_finalize_no_update() {
let sum_agg = bind_aggregate_function(
&FUNCTION_SET_SUM,
vec![expr::column((0, 1), DataType::int64()).into()],
)
.unwrap();
let aggs = [PhysicalAggregateExpression::new(
sum_agg,
[(1, DataType::int64())],
)];
let layout = AggregateLayout::try_new([DataType::utf8()], aggs).unwrap();
let mut collection = AggregateCollection::new(layout, 16);
let mut state = collection.init_append_state();
collection
.append_groups(
&mut state,
&[Array::try_from_iter(["group_a", "group_b"]).unwrap()],
0..2,
)
.unwrap();
let mut ptrs = state.row_pointers().to_vec();
let mut groups = Array::new(&DefaultBufferManager, DataType::utf8(), 2).unwrap();
let mut results = Array::new(&DefaultBufferManager, DataType::int64(), 2).unwrap();
unsafe {
collection
.finalize_groups(&mut ptrs, &mut [&mut groups], &mut [&mut results])
.unwrap();
}
let expected_groups = Array::try_from_iter(["group_a", "group_b"]).unwrap();
let expected_results = Array::try_from_iter([None as Option<i64>, None]).unwrap();
assert_arrays_eq(&expected_groups, &groups);
assert_arrays_eq(&expected_results, &results);
}
#[test]
fn append_groups_finalize_with_update() {
let sum_agg = bind_aggregate_function(
&FUNCTION_SET_SUM,
vec![expr::column((0, 1), DataType::int64()).into()],
)
.unwrap();
let aggs = [PhysicalAggregateExpression::new(
sum_agg,
[(1, DataType::int64())],
)];
let layout = AggregateLayout::try_new([DataType::utf8()], aggs).unwrap();
let mut collection = AggregateCollection::new(layout, 16);
let mut state = collection.init_append_state();
collection
.append_groups(
&mut state,
&[Array::try_from_iter(["group_a", "group_b"]).unwrap()],
0..2,
)
.unwrap();
let ptrs = state.row_pointers();
let mut update_row_ptrs = vec![ptrs[0], ptrs[1], ptrs[0], ptrs[1]];
let values = Array::try_from_iter([1_i64, 2, 3, 4]).unwrap();
unsafe {
collection
.layout
.update_states(
&mut update_row_ptrs,
[AggregateUpdateSelector {
aggregate_idx: 0,
inputs: &[values],
}],
4,
)
.unwrap();
}
let mut ptrs = state.row_pointers().to_vec();
let mut groups = Array::new(&DefaultBufferManager, DataType::utf8(), 2).unwrap();
let mut results = Array::new(&DefaultBufferManager, DataType::int64(), 2).unwrap();
unsafe {
collection
.finalize_groups(&mut ptrs, &mut [&mut groups], &mut [&mut results])
.unwrap();
}
let expected_groups = Array::try_from_iter(["group_a", "group_b"]).unwrap();
let expected_results = Array::try_from_iter([4_i64, 6]).unwrap();
assert_arrays_eq(&expected_groups, &groups);
assert_arrays_eq(&expected_results, &results);
}
}