use std::collections::HashMap;
use rudb_common::{LogicalType, Result};
use rudb_plan::{ColumnBinding, Expr, ExprRef, JoinKind, Node, NodeRef, Plan};
use crate::pass::{Context, Pass};
use crate::walk;
#[derive(Debug, Clone, Copy)]
pub struct EagerAggregation;
impl Pass for EagerAggregation {
fn name(&self) -> &'static str {
"eager_aggregation"
}
fn run(&self, plan: &mut Plan, _context: &Context) -> Result<()> {
push(plan);
Ok(())
}
}
pub fn push(plan: &mut Plan) {
let mut moved = false;
let root = walk::restack(plan, plan.root(), &mut moved, &mut split);
if moved {
plan.set_root(root);
}
}
#[derive(Clone, Copy)]
enum Step {
Filter { predicate: ExprRef },
Join { at: NodeRef, left: bool },
}
fn split(plan: &mut Plan, at: NodeRef) -> Option<NodeRef> {
let Node::Aggregate { input, index, groups, aggregates } = *plan.node(at) else { return None };
let calls = plan.expr_list(aggregates).to_vec();
if calls.is_empty() || !calls.iter().all(|&call| movable(plan, call)) {
return None;
}
let keys = plan.expr_list(groups).to_vec();
let mut read = Vec::new();
for &call in &calls {
walk::columns(plan, call, &mut |binding| read.push(binding));
}
let mut path: Vec<Step> = Vec::new();
let mut here = input;
loop {
match *plan.node(here) {
Node::Filter { input, predicate } => {
path.push(Step::Filter { predicate });
here = input;
}
Node::Join { left, right, kind: JoinKind::Inner, .. } => {
let left_side = walk::outputs(plan, left)?;
let right_side = walk::outputs(plan, right)?;
let within = |side: &[(ColumnBinding, LogicalType)]| {
read.iter().all(|binding| side.iter().any(|(found, _)| found == binding))
};
let (below, beside, side, other) = if within(&left_side) {
(left, right, left_side, right_side)
} else if within(&right_side) {
(right, left, right_side, left_side)
} else {
return None;
};
path.push(Step::Join { at: here, left: below == left });
let scan = matches!(plan.node(beside), Node::Get { .. });
if let Some(kept) =
scan.then(|| worth(plan, &path, &keys, below, &side, &other)).flatten()
{
return Some(rewrite(plan, &path, below, kept, index, &keys, &calls));
}
here = below;
}
_ => return None,
}
}
}
fn movable(plan: &Plan, call: ExprRef) -> bool {
let Expr::Aggregate { name, args, distinct, filter } = *plan.expr(call) else { return false };
let [arg] = plan.expr_list(args) else { return false };
if distinct || filter.is_some() {
return false;
}
match plan.string(name) {
"sum" => matches!(plan.expr_type(*arg), LogicalType::Decimal { .. }),
"min" | "max" => true,
_ => false,
}
}
fn worth(
plan: &Plan,
path: &[Step],
keys: &[ExprRef],
below: NodeRef,
side: &[(ColumnBinding, LogicalType)],
other: &[(ColumnBinding, LogicalType)],
) -> Option<Vec<(ColumnBinding, LogicalType)>> {
if matches!(plan.node(below), Node::Aggregate { .. }) {
return None;
}
let mut wanted = Vec::new();
let mut wide = false;
for &key in keys {
walk::columns(plan, key, &mut |binding| {
wanted.push(binding);
if other.iter().any(|(found, ty)| *found == binding && ty == &LogicalType::Varchar) {
wide = true;
}
});
}
if !wide {
return None;
}
for step in path {
match *step {
Step::Filter { predicate } => {
walk::columns(plan, predicate, &mut |binding| wanted.push(binding));
}
Step::Join { at, .. } => {
let Node::Join { conditions, .. } = *plan.node(at) else { return None };
for &condition in plan.expr_list(conditions) {
walk::columns(plan, condition, &mut |binding| wanted.push(binding));
}
}
}
}
let kept: Vec<(ColumnBinding, LogicalType)> =
side.iter().filter(|(binding, _)| wanted.contains(binding)).cloned().collect();
let narrow = |ty: &LogicalType| {
ty.is_integer() || ty.is_temporal() || matches!(ty, LogicalType::Decimal { .. })
};
(!kept.is_empty() && kept.iter().all(|(_, ty)| narrow(ty))).then_some(kept)
}
fn rewrite(
plan: &mut Plan,
path: &[Step],
below: NodeRef,
kept: Vec<(ColumnBinding, LogicalType)>,
index: u32,
keys: &[ExprRef],
calls: &[ExprRef],
) -> NodeRef {
let staged = walk::fresh_index(plan);
let column = |plan: &mut Plan, at: usize, ty: LogicalType| {
let at = u32::try_from(at).expect("an aggregate with this many expressions cannot bind");
plan.add_expr(Expr::Column(ColumnBinding::new(staged, at)), ty)
};
let mut moved = HashMap::new();
let mut groups = Vec::with_capacity(kept.len());
for (at, (binding, ty)) in kept.iter().enumerate() {
groups.push(plan.add_expr(Expr::Column(*binding), ty.clone()));
moved.insert(*binding, ColumnBinding::new(staged, u32::try_from(at).unwrap_or(u32::MAX)));
}
let groups = plan.add_expr_list(&groups);
let partials = plan.add_expr_list(calls);
let mut built = plan.add_node(Node::Aggregate {
input: below,
index: staged,
groups,
aggregates: partials,
});
for step in path.iter().rev() {
let node = match *step {
Step::Filter { predicate } => {
Node::Filter { input: built, predicate: rebind(plan, predicate, &moved) }
}
Step::Join { at, left } => {
let mut node = plan.node(at).clone();
if let Node::Join { left: l, right: r, conditions, .. } = &mut node {
if left {
*l = built;
} else {
*r = built;
}
let held = plan.expr_list(*conditions).to_vec();
let rebound: Vec<ExprRef> =
held.iter().map(|&condition| rebind(plan, condition, &moved)).collect();
if rebound != held {
*conditions = plan.add_expr_list(&rebound);
}
}
node
}
};
built = plan.add_node(node);
}
let outer_keys: Vec<ExprRef> = keys.iter().map(|&key| rebind(plan, key, &moved)).collect();
let mut outer_calls = Vec::with_capacity(calls.len());
for (at, &call) in calls.iter().enumerate() {
let Expr::Aggregate { name, .. } = *plan.expr(call) else { continue };
let ty = plan.expr_type(call).clone();
let span = plan.expr_span(call);
let arg = column(plan, kept.len() + at, ty.clone());
let args = plan.add_expr_list(&[arg]);
let total = Expr::Aggregate { name, args, distinct: false, filter: None };
outer_calls.push(plan.add_expr_at(total, ty, span));
}
let groups = plan.add_expr_list(&outer_keys);
let aggregates = plan.add_expr_list(&outer_calls);
plan.add_node(Node::Aggregate { input: built, index, groups, aggregates })
}
fn rebind(
plan: &mut Plan,
expr: ExprRef,
moved: &HashMap<ColumnBinding, ColumnBinding>,
) -> ExprRef {
if let Expr::Column(binding) = *plan.expr(expr) {
return match moved.get(&binding) {
Some(&to) => {
let ty = plan.expr_type(expr).clone();
let span = plan.expr_span(expr);
plan.add_expr_at(Expr::Column(to), ty, span)
}
None => expr,
};
}
walk::rebuild(plan, expr, &mut |plan, child| rebind(plan, child, moved))
}
#[cfg(test)]
mod tests {
use rudb_plan::Plan;
use super::push;
fn pushed(text: &str) -> String {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
push(&mut plan);
plan.validate().unwrap_or_else(|error| panic!("{text} did not stay valid: {error}"));
let once = plan.to_string();
push(&mut plan);
assert_eq!(plan.to_string(), once, "a second run moved the plan again");
once
}
const JOINED: &str = concat!(
"Aggregate #3 groups=[#0.1::VARCHAR] aggregates=[sum(#1.1::DECIMAL(15,2))::DECIMAL(38,2)]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT, name::VARCHAR]\n",
" Get memory.main.o AS o #1 [c::BIGINT, price::DECIMAL(15,2)]\n",
);
#[test]
fn a_sum_moves_below_a_join_that_brings_in_a_string_to_group_by() {
assert_eq!(
pushed(JOINED),
concat!(
"Aggregate #3 groups=[#0.1::VARCHAR] aggregates=[sum(#4.1::DECIMAL(38,2))::DECIMAL(38,2)]\n",
" Join INNER on=[(#0.0::BIGINT = #4.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT, name::VARCHAR]\n",
" Aggregate #4 groups=[#1.0::BIGINT] aggregates=[sum(#1.1::DECIMAL(15,2))::DECIMAL(38,2)]\n",
" Get memory.main.o AS o #1 [c::BIGINT, price::DECIMAL(15,2)]\n",
)
);
}
#[test]
fn a_count_or_a_grouping_on_numbers_alone_is_left_where_it_was() {
let counted = JOINED.replace(
"sum(#1.1::DECIMAL(15,2))::DECIMAL(38,2)",
"count(#1.1::DECIMAL(15,2))::BIGINT",
);
assert_eq!(pushed(&counted), counted);
let numbers = JOINED.replace("groups=[#0.1::VARCHAR]", "groups=[#0.0::BIGINT]");
assert_eq!(pushed(&numbers), numbers);
}
#[test]
fn a_side_whose_kept_columns_are_strings_is_passed_over_for_the_one_below_it() {
let text = concat!(
"Aggregate #4 groups=[#0.1::VARCHAR, #2.1::VARCHAR] aggregates=[max(#1.1::DECIMAL(15,2))::DECIMAL(15,2)]\n",
" Join INNER on=[(#0.2::INTEGER = #2.0::INTEGER)::BOOLEAN]\n",
" Get memory.main.n AS n #2 [k::INTEGER, name::VARCHAR]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT, name::VARCHAR, n::INTEGER]\n",
" Filter (#1.1::DECIMAL(15,2) > 1.00::DECIMAL(15,2))::BOOLEAN\n",
" Get memory.main.o AS o #1 [c::BIGINT, price::DECIMAL(15,2)]\n",
);
assert_eq!(
pushed(text),
concat!(
"Aggregate #4 groups=[#0.1::VARCHAR, #2.1::VARCHAR] aggregates=[max(#5.1::DECIMAL(15,2))::DECIMAL(15,2)]\n",
" Join INNER on=[(#0.2::INTEGER = #2.0::INTEGER)::BOOLEAN]\n",
" Get memory.main.n AS n #2 [k::INTEGER, name::VARCHAR]\n",
" Join INNER on=[(#0.0::BIGINT = #5.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT, name::VARCHAR, n::INTEGER]\n",
" Aggregate #5 groups=[#1.0::BIGINT] aggregates=[max(#1.1::DECIMAL(15,2))::DECIMAL(15,2)]\n",
" Filter (#1.1::DECIMAL(15,2) > 1.00::DECIMAL(15,2))::BOOLEAN\n",
" Get memory.main.o AS o #1 [c::BIGINT, price::DECIMAL(15,2)]\n",
)
);
}
#[test]
fn a_join_that_throws_rows_away_keeps_the_sum_above_it() {
let filtered = JOINED.replace(
" Get memory.main.c AS c #0 [k::BIGINT, name::VARCHAR]\n",
" Filter (#0.0::BIGINT > 5::BIGINT)::BOOLEAN\n Get memory.main.c AS c #0 [k::BIGINT, name::VARCHAR]\n",
);
assert_eq!(pushed(&filtered), filtered);
}
}