use gnitz_expr::{order_locators, Lead, OrderLocator, RowRanking};
use std::num::NonZeroI64;
use gnitz_wire::{ReadSink, RowsCut, SinkKind};
use super::aggregate::AdhocFold;
use super::map::MapPlan;
use crate::repr::Batch;
use crate::schema::SchemaDescriptor;
pub struct SinkPlan {
map: Option<MapPlan>,
kind: Kind,
}
enum Kind {
Rows { keeper: Batch, cut: Option<Cut> },
Fold {
fold: Box<AdhocFold>,
mapped: Option<Batch>,
},
}
enum Cut {
Prefix { window: NonZeroI64, summed: i64 },
TopK {
order: Vec<OrderLocator>,
window: NonZeroI64,
summed: i64,
bound: Option<Lead>,
},
}
impl SinkPlan {
pub fn from_wire(src_schema: &SchemaDescriptor, sink: &ReadSink, group_cap: usize) -> Result<Self, String> {
let map = sink
.map
.as_ref()
.map(|m| MapPlan::from_compute_map(src_schema, m))
.transpose()
.map_err(|e| format!("scan_spec map: {e}"))?;
let sink_in = map.as_ref().map_or(*src_schema, |m| *m.out_schema());
let kind = match &sink.kind {
SinkKind::Fold(agg) => Kind::Fold {
fold: Box::new(AdhocFold::new(&sink_in, agg, group_cap)?),
mapped: map.as_ref().map(|m| Batch::empty_with_schema(m.out_schema())),
},
SinkKind::Rows { cut } => {
let cut = match cut {
None => None,
Some(RowsCut { k, order }) => {
for k in order {
sink_in.wire_col("scan_spec: order key column", k.col as u32)?;
}
let window = NonZeroI64::try_from(*k).unwrap_or(NonZeroI64::MAX);
Some(match order.is_empty() {
true => Cut::Prefix { window, summed: 0 },
false => Cut::TopK {
order: order_locators(order, &sink_in, true),
window,
summed: 0,
bound: None,
},
})
}
};
Kind::Rows {
keeper: Batch::empty_with_schema(&sink_in),
cut,
}
}
};
Ok(SinkPlan { map, kind })
}
pub fn output_schema(&self) -> &SchemaDescriptor {
match &self.kind {
Kind::Rows { keeper, .. } => keeper.schema(),
Kind::Fold { fold, .. } => fold.output_schema(),
}
}
pub fn first_drain(&self, chunk_rows: usize) -> usize {
match &self.kind {
Kind::Rows {
cut: Some(Cut::Prefix { window, .. }), ..
} => (window.get() as usize).min(chunk_rows),
_ => chunk_rows,
}
}
pub fn push(&mut self, chunk: &Batch, ranges: &mut Vec<(usize, usize)>) -> Result<bool, String> {
let mb = chunk.as_mem_batch();
match &mut self.kind {
Kind::Rows { keeper, cut: None } => {
append_survivors(self.map.as_mut(), chunk, keeper, ranges);
Ok(false)
}
Kind::Rows {
keeper,
cut: Some(Cut::Prefix { window, summed }),
} => {
let window = window.get();
for i in 0..ranges.len() {
let (s, e) = ranges[i];
let range_sum = mb.sum_weights(s, e);
if *summed + range_sum < window {
*summed += range_sum;
continue;
}
let mut end = s;
while *summed < window && end < e {
*summed += mb.get_weight(end);
end += 1;
}
ranges[i].1 = end;
ranges.truncate(i + 1);
break;
}
append_survivors(self.map.as_mut(), chunk, keeper, ranges);
Ok(*summed >= window)
}
Kind::Rows {
keeper,
cut: Some(Cut::TopK { order, window, summed, bound }),
} => {
*summed = ranges
.iter()
.fold(*summed, |a, &(s, e)| a.wrapping_add(mb.sum_weights(s, e)));
append_survivors(self.map.as_mut(), chunk, keeper, ranges);
if *summed > window.get().saturating_mul(2) {
*summed = topk_keep(keeper, order, *window, bound);
}
Ok(false)
}
Kind::Fold { fold, mapped } => {
match (&mut self.map, mapped) {
(Some(plan), Some(dst)) => {
dst.clear();
plan.append_map_ranges(chunk, dst, ranges);
fold.fold_ranges(dst, &[(0, dst.len())])?;
}
_ => fold.fold_ranges(chunk, ranges)?,
}
Ok(false)
}
}
}
pub fn finish(self) -> Batch {
match self.kind {
Kind::Rows { mut keeper, cut } => {
if let Some(Cut::TopK { order, window, summed, mut bound }) = cut {
if summed > window.get() {
topk_keep(&mut keeper, &order, window, &mut bound);
}
}
keeper
}
Kind::Fold { fold, .. } => fold.finish(),
}
}
}
fn append_survivors(map: Option<&mut MapPlan>, chunk: &Batch, keeper: &mut Batch, ranges: &[(usize, usize)]) {
match map {
None => keeper.append_ranges(&chunk.as_mem_batch(), ranges),
Some(p) => p.append_map_ranges(chunk, keeper, ranges),
}
}
fn topk_keep(keeper: &mut Batch, order: &[OrderLocator], window: NonZeroI64, bound: &mut Option<Lead>) -> i64 {
if keeper.is_empty() {
return 0;
}
let window = window.get();
let mut ranking = RowRanking::new(order, &*keeper);
if let Some(bound) = *bound {
ranking.drop_above(bound);
}
ranking.keep_smallest(window as usize);
*bound = ranking.max_lead();
let mut acc: i64 = ranking.rows().map(|r| keeper.get_weight(r as usize)).sum();
let perm = if acc > window {
let mut perm = ranking.sorted();
acc = 0;
let cut = perm.iter().position(|&r| {
acc += keeper.get_weight(r as usize);
acc >= window
});
perm.truncate(cut.map_or(perm.len(), |i| i + 1));
perm
} else {
ranking.rows().collect()
};
*keeper = keeper.indexed_rows(&perm);
acc
}