use std::collections::HashMap;
use rudb_common::{LogicalType, Result, Value};
use rudb_plan::{ColumnBinding, CompareOp, Expr, ExprRef, JoinKind, Node, NodeRef, Plan};
use crate::estimate::{self, Facts};
use crate::link::Linked;
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, context.facts(), context.links());
Ok(())
}
}
pub fn push(plan: &mut Plan, stats: &Facts, links: &[Linked]) {
let mut moved = false;
let root =
walk::restack(plan, plan.root(), &mut moved, &mut |plan, at| split(plan, at, stats, links));
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, stats: &Facts, links: &[Linked]) -> 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));
}
if read.is_empty() {
return None;
}
let mut path: Vec<Step> = Vec::new();
let mut padded = false;
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: kind @ (JoinKind::Inner | JoinKind::Left), .. } => {
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 });
padded |= kind == JoinKind::Left && below == right;
let scan = matches!(plan.node(beside), Node::Get { .. });
if let Some(kept) =
scan.then(|| worth(plan, &path, &keys, below, &side, &other, stats)).flatten()
{
let joins =
path.iter().filter(|step| matches!(step, Step::Join { .. })).count();
let alone = joins == 1
&& covers(plan, here, &kept, &other)
&& keyed(plan, here, beside, &keys, links);
let staged = Staged { below, kept, padded, alone };
return Some(rewrite(plan, &path, staged, 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 };
if distinct || filter.is_some() {
return false;
}
let name = plan.string(name);
let args = plan.expr_list(args);
if name == "count_star" {
return args.is_empty() && plan.expr_type(call) == &LogicalType::BigInt;
}
let [arg] = args else { return false };
match name {
"sum" => matches!(plan.expr_type(*arg), LogicalType::Decimal { .. }),
"count" => plan.expr_type(call) == &LogicalType::BigInt,
"min" | "max" => true,
_ => false,
}
}
fn worth(
plan: &Plan,
path: &[Step],
keys: &[ExprRef],
below: NodeRef,
side: &[(ColumnBinding, LogicalType)],
other: &[(ColumnBinding, LogicalType)],
stats: &Facts,
) -> 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;
}
});
}
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 { .. })
};
let fits = !kept.is_empty() && kept.iter().all(|(_, ty)| narrow(ty));
(fits && (wide || shrinks(plan, below, &kept, stats))).then_some(kept)
}
const SHRINK: u64 = 4;
fn shrinks(
plan: &Plan,
below: NodeRef,
kept: &[(ColumnBinding, LogicalType)],
stats: &Facts,
) -> bool {
let Some(&rows) = estimate::unfiltered(plan, below, stats).value() else { return false };
let mut groups: u64 = 1;
for (binding, _) in kept {
let Some(&distinct) = estimate::stated(plan, *binding, stats).value() else {
return false;
};
groups = groups.saturating_mul(distinct);
}
rows > 0 && groups.min(rows).saturating_mul(SHRINK) <= rows
}
struct Staged {
below: NodeRef,
kept: Vec<(ColumnBinding, LogicalType)>,
padded: bool,
alone: bool,
}
fn covers(
plan: &Plan,
at: NodeRef,
kept: &[(ColumnBinding, LogicalType)],
other: &[(ColumnBinding, LogicalType)],
) -> bool {
let Node::Join { conditions, .. } = *plan.node(at) else { return false };
let column = |expr: ExprRef| match *plan.expr(expr) {
Expr::Column(binding) => Some(binding),
_ => None,
};
let ours = |binding: ColumnBinding| kept.iter().any(|(found, _)| *found == binding);
let theirs = |binding: ColumnBinding| other.iter().any(|(found, _)| *found == binding);
let mut met = Vec::new();
for &condition in plan.expr_list(conditions) {
let Expr::Compare { op: CompareOp::Equal, left, right } = *plan.expr(condition) else {
return false;
};
let (Some(left), Some(right)) = (column(left), column(right)) else { return false };
if ours(left) && theirs(right) {
met.push(left);
} else if theirs(left) && ours(right) {
met.push(right);
} else {
return false;
}
}
kept.iter().all(|(binding, _)| met.contains(binding))
}
fn keyed(plan: &Plan, join: NodeRef, beside: NodeRef, keys: &[ExprRef], links: &[Linked]) -> bool {
let Node::Join { kind, conditions, .. } = *plan.node(join) else { return false };
let Node::Get { table, columns, index, .. } = *plan.node(beside) else { return false };
let table = plan.string(table);
let fields = plan.field_list(columns);
let compared = |binding: ColumnBinding| {
plan.expr_list(conditions).iter().any(|&condition| {
matches!(*plan.expr(condition), Expr::Compare { op: CompareOp::Equal, left, right }
if [left, right].iter().any(|&side| *plan.expr(side) == Expr::Column(binding)))
})
};
keys.iter().any(|&key| {
let Expr::Column(binding) = *plan.expr(key) else { return false };
if binding.table != index {
return false;
}
let Some(field) = fields.get(binding.column as usize) else { return false };
let never_null = field.not_null
|| estimate::never_null(plan, beside, binding)
|| (kind == JoinKind::Inner && compared(binding));
never_null
&& links.iter().any(|link| {
link.unique
&& link.second.is_none()
&& link.parent.eq_ignore_ascii_case(table)
&& link.parent_column.eq_ignore_ascii_case(&field.name)
})
})
}
fn rewrite(
plan: &mut Plan,
path: &[Step],
staged_at: Staged,
index: u32,
keys: &[ExprRef],
calls: &[ExprRef],
) -> NodeRef {
let Staged { below, kept, padded, alone } = staged_at;
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 mut arg = column(plan, kept.len() + at, ty.clone());
let (name, nothing) = match plan.string(name) {
"count" => (plan.intern("sum"), Some(0)),
"count_star" => (plan.intern("sum"), Some(1)),
_ => (name, None),
};
if let Some(nothing) = nothing.filter(|_| padded) {
let nothing = plan.add_constant(Value::BigInt(nothing));
let args = plan.add_expr_list(&[arg, nothing]);
let coalesce = Expr::Function { name: plan.intern("coalesce"), args };
arg = plan.add_expr_at(coalesce, ty.clone(), span);
}
if alone {
outer_calls.push(arg);
continue;
}
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));
}
if alone {
let exprs: Vec<ExprRef> = outer_keys.into_iter().chain(outer_calls).collect();
let names: Vec<_> =
(0..exprs.len()).map(|position| plan.intern(&format!("column{position}"))).collect();
let exprs = plan.add_expr_list(&exprs);
let names = plan.add_name_list(&names);
return plan.add_node(Node::Project { input: built, index, exprs, names });
}
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_common::Provenance;
use rudb_plan::Plan;
use super::push;
use crate::estimate::Facts;
use crate::link::Linked;
fn pushed(text: &str) -> String {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
push(&mut plan, &Facts::new(), &[]);
plan.validate().unwrap_or_else(|error| panic!("{text} did not stay valid: {error}"));
let once = plan.to_string();
push(&mut plan, &Facts::new(), &[]);
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_moves_as_a_sum_of_counts_that_keeps_its_type() {
let counted = JOINED.replace(
"sum(#1.1::DECIMAL(15,2))::DECIMAL(38,2)",
"count(#1.1::DECIMAL(15,2))::BIGINT, count_star()::BIGINT",
);
assert_eq!(
pushed(&counted),
concat!(
"Aggregate #3 groups=[#0.1::VARCHAR] aggregates=[sum(#4.1::BIGINT)::BIGINT, sum(#4.2::BIGINT)::BIGINT]\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=[count(#1.1::DECIMAL(15,2))::BIGINT, count_star()::BIGINT]\n",
" Get memory.main.o AS o #1 [c::BIGINT, price::DECIMAL(15,2)]\n",
)
);
}
#[test]
fn a_lone_count_star_or_a_grouping_on_numbers_alone_is_left_where_it_was() {
let star =
JOINED.replace("sum(#1.1::DECIMAL(15,2))::DECIMAL(38,2)", "count_star()::BIGINT");
assert_eq!(pushed(&star), star);
let numbers = JOINED.replace("groups=[#0.1::VARCHAR]", "groups=[#0.0::BIGINT]");
assert_eq!(pushed(&numbers), numbers);
}
const PADDED: &str = concat!(
"Aggregate #3 groups=[#0.0::BIGINT] aggregates=[count(#1.1::BIGINT)::BIGINT, count_star()::BIGINT]\n",
" Join LEFT on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT]\n",
" Filter (#1.1::BIGINT > 5::BIGINT)::BOOLEAN\n",
" Get memory.main.o AS o #1 [c::BIGINT, key::BIGINT]\n",
);
fn counted(rows: u64, distinct: u64) -> Facts {
let mut facts = Facts::new();
facts.record("memory", "main", "o", rows);
facts.record_distinct("memory", "main", "o", "c", distinct, Provenance::Dictionary);
facts
}
#[test]
fn a_count_below_the_padded_side_of_a_left_join_reads_coalesce_above_it() {
let stats = counted(1_500_000, 100_000);
let mut plan = Plan::parse(PADDED).expect("parses");
push(&mut plan, &stats, &[]);
plan.validate().expect("valid");
assert_eq!(
plan.to_string(),
concat!(
"Aggregate #3 groups=[#0.0::BIGINT] aggregates=[sum(coalesce(#4.1::BIGINT, 0::BIGINT)::BIGINT)::BIGINT, sum(coalesce(#4.2::BIGINT, 1::BIGINT)::BIGINT)::BIGINT]\n",
" Join LEFT on=[(#0.0::BIGINT = #4.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT]\n",
" Aggregate #4 groups=[#1.0::BIGINT] aggregates=[count(#1.1::BIGINT)::BIGINT, count_star()::BIGINT]\n",
" Filter (#1.1::BIGINT > 5::BIGINT)::BOOLEAN\n",
" Get memory.main.o AS o #1 [c::BIGINT, key::BIGINT]\n",
)
);
let once = plan.to_string();
push(&mut plan, &stats, &[]);
assert_eq!(plan.to_string(), once, "a second run moved the plan again");
}
#[test]
fn a_group_key_a_link_says_is_unique_makes_the_top_a_projection() {
let stats = counted(1_500_000, 100_000);
let link = |unique: bool| Linked { unique, ..Linked::declared("o", "c", "c", "k") };
let inner = PADDED.replace("Join LEFT", "Join INNER");
let mut plan = Plan::parse(&inner).expect("parses");
push(&mut plan, &stats, &[link(true)]);
plan.validate().expect("valid");
assert_eq!(
plan.to_string(),
concat!(
"Project #3 [#0.0::BIGINT AS column0, #4.1::BIGINT AS column1, #4.2::BIGINT AS column2]\n",
" Join INNER on=[(#0.0::BIGINT = #4.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT]\n",
" Aggregate #4 groups=[#1.0::BIGINT] aggregates=[count(#1.1::BIGINT)::BIGINT, count_star()::BIGINT]\n",
" Filter (#1.1::BIGINT > 5::BIGINT)::BOOLEAN\n",
" Get memory.main.o AS o #1 [c::BIGINT, key::BIGINT]\n",
)
);
for (text, built) in [(inner.as_str(), false), (PADDED, true)] {
let mut plan = Plan::parse(text).expect("parses");
push(&mut plan, &stats, &[link(built)]);
assert!(plan.to_string().starts_with("Aggregate #3"), "{plan}");
}
}
#[test]
fn a_grouping_on_numbers_moves_only_where_the_file_says_it_shrinks_by_four() {
for (rows, distinct, moves) in
[(1_500_000, 100_000, true), (400, 100, true), (400, 101, false)]
{
let mut plan = Plan::parse(PADDED).expect("parses");
push(&mut plan, &counted(rows, distinct), &[]);
assert_eq!(plan.to_string() != PADDED, moves, "{rows} rows over {distinct} values");
}
assert_eq!(pushed(PADDED), PADDED, "nothing measured");
}
#[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);
}
}