use gnitz_expr::{cmp_order_keys, order_locators, OrderLocator};
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>,
},
}
struct Cut {
order: Vec<OrderLocator>,
window: NonZeroI64,
summed: i64,
}
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 }) => {
sink_in.check_cols(order.iter().map(|k| ("scan_spec: order key column", k.col as u32)))?;
Some(Cut {
order: order_locators(order, &sink_in),
window: NonZeroI64::try_from(*k).unwrap_or(NonZeroI64::MAX),
summed: 0,
})
}
};
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), .. } if cut.order.is_empty() => (cut.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) } if cut.order.is_empty() => {
let window = cut.window.get();
for i in 0..ranges.len() {
let (s, e) = ranges[i];
let range_sum = mb.sum_weights(s, e);
if cut.summed + range_sum < window {
cut.summed += range_sum;
continue;
}
let mut end = s;
while cut.summed < window && end < e {
cut.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(cut.summed >= window)
}
Kind::Rows { keeper, cut: Some(cut) } => {
cut.summed = ranges
.iter()
.fold(cut.summed, |a, &(s, e)| a.wrapping_add(mb.sum_weights(s, e)));
append_survivors(self.map.as_mut(), chunk, keeper, ranges);
if cut.summed > cut.window.get().saturating_mul(2) {
cut.summed = topk_keep(keeper, &cut.order, cut.window);
}
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) = cut.filter(|c| !c.order.is_empty() && c.summed > c.window.get()) {
topk_keep(&mut keeper, &cut.order, cut.window);
}
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) -> i64 {
if keeper.is_empty() {
return 0;
}
let window = window.get();
let mut perm: Vec<u32> = (0..keeper.len() as u32).collect();
let cmp = |a: &u32, b: &u32| {
let (ra, rb) = (*a as usize, *b as usize);
cmp_order_keys(order, &*keeper, ra, &*keeper, rb)
};
let k = (window as usize).min(perm.len());
if k < perm.len() {
perm.select_nth_unstable_by(k - 1, cmp);
perm.truncate(k);
}
let mut acc: i64 = perm.iter().map(|&r| keeper.get_weight(r as usize)).sum();
if acc > window {
perm.sort_unstable_by(cmp);
acc = 0;
let cut = perm
.iter()
.position(|&r| {
acc += keeper.get_weight(r as usize);
acc >= window
})
.map_or(perm.len(), |i| i + 1);
perm.truncate(cut);
}
*keeper = keeper.indexed_rows(&perm);
acc
}