use gnitz_wire::AggReadSpec;
use super::agg::{Accumulator, GroupedState};
use super::emit::emit_reduce_row;
use super::shape::ReduceShape;
use crate::algebra::group_key::{ground_pk, GroupNumbers, GroupOutKey, IdentityLoop};
use crate::repr::{Batch, MemBatch};
use crate::schema::{SchemaDescriptor, SchemaFacts};
pub(crate) struct AdhocFold {
shape: ReduceShape,
groups: Batch,
global: Vec<Accumulator>,
states: Vec<GroupedState>,
numbers: GroupNumbers,
ord: Vec<u32>,
group_cap: usize,
}
impl AdhocFold {
pub(crate) fn new(src_schema: &SchemaDescriptor, agg: &AggReadSpec, group_cap: usize) -> Result<Self, String> {
let refuse = |e| format!("scan_spec fold: {e}");
let (key, prefix) =
GroupOutKey::new(src_schema, &agg.group_cols, agg.group_cols.iter().copied()).map_err(refuse)?;
let mut groups = Batch::empty_with_schema(&prefix.finish().map_err(refuse)?);
let shape = ReduceShape::new(src_schema, key, prefix, &agg.aggs).map_err(refuse)?;
let mut global = Vec::new();
if shape.key.is_global() {
groups.push_key_row(ground_pk().bytes(), 1);
global.extend_from_slice(&shape.acc_template);
}
Ok(AdhocFold {
states: shape.acc_template.iter().map(Accumulator::grouped).collect(),
global,
groups,
shape,
numbers: GroupNumbers::default(),
ord: Vec::new(),
group_cap,
})
}
pub(crate) fn output_schema(&self) -> &SchemaDescriptor {
&self.shape.output_schema
}
pub(crate) fn fold_ranges(&mut self, chunk: &Batch, ranges: &[(usize, usize)]) -> Result<(), String> {
let Self {
shape,
groups,
global,
states,
numbers,
ord,
group_cap,
} = self;
let mb = chunk.as_mem_batch();
if shape.key.is_global() {
for &(s, e) in ranges {
Accumulator::fold_rows(global, &mb, s..e, true);
}
return Ok(());
}
ord.clear();
let assign = AssignGroups {
shape,
groups,
numbers,
ord,
group_cap: *group_cap,
mb: &mb,
ranges,
};
shape.key.with_identity(&mb, assign)?;
for (acc, state) in shape.acc_template.iter().zip(states) {
state.resize(groups.count, acc);
acc.fold_grouped(&mb, ranges, ord, state);
}
Ok(())
}
pub(crate) fn finish(mut self) -> Batch {
let gs = self.groups.schema();
let carried = gs.payload_locators();
let mut output = Batch::with_capacity(&self.shape.output_schema, self.groups.count);
let groups_mb = self.groups.as_mem_batch();
let grouped = !self.shape.key.is_global();
let mut accs = match grouped {
true => self.shape.acc_template.clone(),
false => self.global,
};
for ord in 0..self.groups.count {
if grouped {
for (acc, state) in accs.iter_mut().zip(&mut self.states) {
state.take(ord, acc);
}
}
emit_reduce_row(
&mut output,
Some((&groups_mb, ord, &carried)),
groups_mb.get_pk_bytes(ord),
&accs,
);
}
output
}
}
struct AssignGroups<'a> {
shape: &'a ReduceShape,
groups: &'a mut Batch,
numbers: &'a mut GroupNumbers,
ord: &'a mut Vec<u32>,
group_cap: usize,
mb: &'a MemBatch<'a>,
ranges: &'a [(usize, usize)],
}
impl IdentityLoop for AssignGroups<'_> {
type Out = Result<(), String>;
fn run(self, identity: impl Fn(usize) -> u128) -> Self::Out {
let AssignGroups {
shape,
groups,
numbers,
ord,
group_cap,
mb,
ranges,
} = self;
for row in ranges.iter().flat_map(|&(s, e)| s..e) {
debug_assert!(
mb.get_weight(row) > 0,
"adhoc fold: scan cursor must deliver positive weights"
);
let g = numbers.ordinal(identity(row), || {
let g = groups.count;
if g >= group_cap {
return Err(format!(
"GROUP BY exceeds {group_cap} distinct groups for ad-hoc execution; \
CREATE VIEW to maintain this aggregation incrementally"
));
}
emit_reduce_row(
groups,
Some((mb, row, shape.key.carried())),
shape.key.out_pk(mb, row).bytes(),
&[],
);
Ok(g as u32)
})?;
ord.push(g);
}
Ok(())
}
}