use rudb_common::{LogicalType, Result};
use rudb_plan::{ColumnBinding, Expr, ExprRef, Node, NodeRef, Plan};
use crate::pass::{Context, Pass};
use crate::walk;
#[derive(Debug, Clone, Copy)]
pub struct CommonAggregate;
impl Pass for CommonAggregate {
fn name(&self) -> &'static str {
"common_aggregate"
}
fn run(&self, plan: &mut Plan, _context: &Context) -> Result<()> {
share(plan);
Ok(())
}
}
pub fn share(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 Answer {
Kept(usize),
Mean { total: usize, count: usize },
}
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();
let mut kept: Vec<ExprRef> = Vec::with_capacity(calls.len());
let mut answers = Vec::with_capacity(calls.len());
let mut shared = false;
for &call in &calls {
let Some((arg, total)) = mean(plan, call).and_then(|arg| {
let total = calls.iter().position(|&other| summed(plan, other, arg))?;
Some((arg, total))
}) else {
answers.push(None);
kept.push(call);
continue;
};
shared = true;
answers.push(Some((arg, total)));
}
if !shared {
return None;
}
let mut placed = Vec::with_capacity(calls.len());
let mut counts: Vec<(ExprRef, usize)> = Vec::new();
for (&call, answer) in calls.iter().zip(&answers) {
placed.push(match *answer {
None => Answer::Kept(kept.iter().position(|&held| held == call)?),
Some((arg, total)) => {
let total = kept.iter().position(|&held| held == calls[total])?;
let count = match counts.iter().find(|(held, _)| walk::same(plan, *held, arg)) {
Some(&(_, count)) => count,
None => {
let count = kept.len() + counts.len();
counts.push((arg, count));
count
}
};
Answer::Mean { total, count }
}
});
}
let mut built = kept.clone();
for &(arg, _) in &counts {
built.push(count(plan, arg));
}
let staged = walk::fresh_index(plan);
let keys = plan.expr_list(groups).to_vec();
let aggregates = plan.add_expr_list(&built);
let inner = plan.add_node(Node::Aggregate { input, index: staged, groups, aggregates });
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 projected = Vec::with_capacity(keys.len() + calls.len());
for (at, &key) in keys.iter().enumerate() {
let ty = plan.expr_type(key).clone();
projected.push(column(plan, at, ty));
}
for (&call, answer) in calls.iter().zip(&placed) {
let ty = plan.expr_type(call).clone();
let expr = match *answer {
Answer::Kept(at) => column(plan, keys.len() + at, ty),
Answer::Mean { total, count } => {
let summed = plan.expr_type(kept[total]).clone();
let total = column(plan, keys.len() + total, summed);
let count = column(plan, keys.len() + count, LogicalType::BigInt);
let name = plan.intern("__rudb_mean");
let args = plan.add_expr_list(&[total, count]);
let span = plan.expr_span(call);
plan.add_expr_at(Expr::Function { name, args }, ty, span)
}
};
projected.push(expr);
}
let names: Vec<_> =
(0..projected.len()).map(|position| plan.intern(&format!("column{position}"))).collect();
let exprs = plan.add_expr_list(&projected);
let names = plan.add_name_list(&names);
Some(plan.add_node(Node::Project { input: inner, index, exprs, names }))
}
fn mean(plan: &Plan, call: ExprRef) -> Option<ExprRef> {
let arg = plain(plan, call, "avg")?;
if plan.expr_type(call) != &LogicalType::Double {
return None;
}
let fits = match plan.expr_type(arg) {
LogicalType::Decimal { width, .. } => *width <= 18,
ty => ty.is_integer() && !matches!(ty, LogicalType::HugeInt | LogicalType::UHugeInt),
};
fits.then_some(arg)
}
fn summed(plan: &Plan, call: ExprRef, arg: ExprRef) -> bool {
let Some(summing) = plain(plan, call, "sum") else { return false };
let total = match (plan.expr_type(call), plan.expr_type(arg)) {
(LogicalType::HugeInt, from) => from.is_integer(),
(LogicalType::Decimal { scale, .. }, LogicalType::Decimal { scale: from, .. }) => {
scale == from
}
_ => false,
};
total && walk::same(plan, summing, arg)
}
fn plain(plan: &Plan, call: ExprRef, name: &str) -> Option<ExprRef> {
let Expr::Aggregate { name: called, args, distinct, filter } = *plan.expr(call) else {
return None;
};
let [arg] = plan.expr_list(args) else { return None };
(!distinct && filter.is_none() && plan.string(called) == name).then_some(*arg)
}
fn count(plan: &mut Plan, arg: ExprRef) -> ExprRef {
let name = plan.intern("count");
let args = plan.add_expr_list(&[arg]);
plan.add_expr(
Expr::Aggregate { name, args, distinct: false, filter: None },
LogicalType::BigInt,
)
}
#[cfg(test)]
mod tests {
use rudb_plan::Plan;
use super::share;
fn shared(text: &str) -> String {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
share(&mut plan);
plan.validate().unwrap_or_else(|error| panic!("{text} did not stay valid: {error}"));
let once = plan.to_string();
share(&mut plan);
assert_eq!(plan.to_string(), once, "a second run changed the plan again");
once
}
const BOTH: &str = concat!(
"Aggregate #1 groups=[#0.0::VARCHAR] aggregates=[avg(#0.1::DECIMAL(15,2))::DOUBLE, ",
"sum(#0.1::DECIMAL(15,2))::DECIMAL(38,2), avg(#0.2::INTEGER)::DOUBLE]\n",
" Get memory.main.t AS t #0 [g::VARCHAR, price::DECIMAL(15,2), n::INTEGER]\n",
);
#[test]
fn an_average_beside_a_sum_of_the_same_column_is_read_off_it() {
assert_eq!(
shared(BOTH),
concat!(
"Project #1 [#2.0::VARCHAR AS column0, __rudb_mean(#2.1::DECIMAL(38,2), ",
"#2.3::BIGINT)::DOUBLE AS column1, #2.1::DECIMAL(38,2) AS column2, ",
"#2.2::DOUBLE AS column3]\n",
" Aggregate #2 groups=[#0.0::VARCHAR] aggregates=[sum(#0.1::DECIMAL(15,2))::DECIMAL(38,2), ",
"avg(#0.2::INTEGER)::DOUBLE, count(#0.1::DECIMAL(15,2))::BIGINT]\n",
" Get memory.main.t AS t #0 [g::VARCHAR, price::DECIMAL(15,2), n::INTEGER]\n",
)
);
}
#[test]
fn an_average_with_no_sum_beside_it_or_a_filter_is_left_alone() {
let alone = BOTH.replace("sum(#0.1::DECIMAL(15,2))", "sum(#0.2::INTEGER)");
let alone = alone.replace("::DECIMAL(38,2), avg(#0.2", "::HUGEINT, max(#0.2");
assert_eq!(shared(&alone), alone);
let distinct =
BOTH.replace("avg(#0.1::DECIMAL(15,2))", "avg(DISTINCT #0.1::DECIMAL(15,2))");
assert_eq!(shared(&distinct), distinct);
}
#[test]
fn a_wide_decimal_is_left_alone() {
let wide = BOTH.replace("DECIMAL(15,2)", "DECIMAL(38,2)");
assert_eq!(shared(&wide), wide);
}
}