use rudb_common::{Class, LogicalType, Result, Stat, Value};
use rudb_plan::{ColumnBinding, CompareOp, Expr, ExprRef, JoinKind, Node, NodeRef, Plan, Slice};
use crate::columns;
use crate::estimate::{self, Facts, Key};
use crate::fromkey::cast;
use crate::pass::{Context, Pass};
use crate::walk;
#[derive(Debug, Clone, Copy)]
pub struct RowsAreGroups;
impl Pass for RowsAreGroups {
fn name(&self) -> &'static str {
"rows_are_groups"
}
fn run(&self, plan: &mut Plan, context: &Context) -> Result<()> {
project_all(plan, context.facts());
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
pub struct JoinedRowsAreGroups;
impl Pass for JoinedRowsAreGroups {
fn name(&self) -> &'static str {
"joined_rows_are_groups"
}
fn run(&self, plan: &mut Plan, context: &Context) -> Result<()> {
if project_all(plan, context.facts()) {
columns::prune(plan);
columns::forward(plan);
}
Ok(())
}
}
pub fn project_all(plan: &mut Plan, stats: &Facts) -> bool {
let mut moved = false;
let root =
walk::restack(plan, plan.root(), &mut moved, &mut |plan, at| project(plan, at, stats));
if moved {
plan.set_root(root);
}
moved
}
fn project(plan: &mut Plan, at: NodeRef, stats: &Facts) -> Option<NodeRef> {
let Node::Aggregate { input, index, groups, aggregates } = *plan.node(at) else { return None };
let keys = plan.expr_list(groups).to_vec();
let calls = plan.expr_list(aggregates).to_vec();
if !keys.iter().any(|&key| distinct(plan, input, key, stats)) {
return None;
}
let mut outputs = keys;
for &call in &calls {
outputs.push(one_row(plan, call)?);
}
let names: Vec<_> =
(0..outputs.len()).map(|position| plan.intern(&format!("column{position}"))).collect();
let exprs = plan.add_expr_list(&outputs);
let names = plan.add_name_list(&names);
Some(plan.add_node(Node::Project { input, index, exprs, names }))
}
fn distinct(plan: &Plan, input: NodeRef, key: ExprRef, stats: &Facts) -> bool {
let Expr::Column(outer) = *plan.expr(key) else { return false };
unique(plan, input, outer, stats)
}
fn unique(plan: &Plan, at: NodeRef, binding: ColumnBinding, stats: &Facts) -> bool {
match *plan.node(at) {
Node::Get { index, .. } if index == binding.table => counted(plan, at, binding, stats),
Node::Filter { input, .. } => unique(plan, input, binding, stats),
Node::Project { input, index, exprs, .. } if index == binding.table => {
let Some(&expr) = plan.expr_list(exprs).get(binding.column as usize) else {
return false;
};
let Expr::Column(inner) = *plan.expr(expr) else { return false };
unique(plan, input, inner, stats)
}
Node::Aggregate { index, groups, .. } if index == binding.table => {
binding.column == 0 && plan.expr_list(groups).len() == 1
}
Node::Join { left, right, kind, conditions, .. } => {
let from_left = produces(plan, left, binding);
match kind {
JoinKind::Semi | JoinKind::Anti | JoinKind::Mark | JoinKind::Single => {
from_left && unique(plan, left, binding, stats)
}
JoinKind::Inner | JoinKind::Left => {
let (mine, other) = if from_left { (left, right) } else { (right, left) };
if !from_left && kind == JoinKind::Left {
return false;
}
(from_left || produces(plan, right, binding))
&& unique(plan, mine, binding, stats)
&& once(plan, conditions, other, stats)
}
_ => false,
}
}
_ => false,
}
}
fn counted(plan: &Plan, at: NodeRef, binding: ColumnBinding, stats: &Facts) -> bool {
let Node::Get { catalog, schema, table, columns, .. } = *plan.node(at) else {
return false;
};
let Some(field) = plan.field_list(columns).get(binding.column as usize) else {
return false;
};
let (catalog, schema, table) = (plan.string(catalog), plan.string(schema), plan.string(table));
let exact = |stat: Stat<u64>| match stat {
Stat::Known { value, class: Class::Exact, .. } => Some(value),
_ => None,
};
let rows = exact(stats.get(&Key::Rows { catalog, schema, table }));
let values = exact(stats.get(&Key::Distinct { catalog, schema, table, column: &field.name }));
let never_null = field.not_null || estimate::never_null(plan, at, binding);
never_null && rows.is_some() && rows == values
}
fn produces(plan: &Plan, at: NodeRef, binding: ColumnBinding) -> bool {
walk::outputs(plan, at)
.is_some_and(|columns| columns.iter().any(|(bound, _)| *bound == binding))
}
fn once(plan: &Plan, conditions: Slice, other: NodeRef, stats: &Facts) -> bool {
let column = |expr: ExprRef| match *plan.expr(expr) {
Expr::Column(binding) => Some(binding),
Expr::Cast { input, try_cast: false }
if plan.expr_type(expr).is_integer() && plan.expr_type(input).is_integer() =>
{
match *plan.expr(input) {
Expr::Column(binding) => Some(binding),
_ => None,
}
}
_ => None,
};
plan.expr_list(conditions).iter().any(|&condition| {
let Expr::Compare { op: CompareOp::Equal, left, right } = *plan.expr(condition) else {
return false;
};
[left, right]
.into_iter()
.filter_map(column)
.any(|binding| produces(plan, other, binding) && unique(plan, other, binding, stats))
})
}
fn one_row(plan: &mut Plan, call: ExprRef) -> Option<ExprRef> {
let Expr::Aggregate { name, args, distinct, filter } = *plan.expr(call) else { return None };
if distinct || filter.is_some() {
return None;
}
let want = plan.expr_type(call).clone();
let span = plan.expr_span(call);
let written = match (plan.string(name), plan.expr_list(args)) {
("count_star", []) => plan.add_constant(Value::BigInt(1)),
("min" | "max", &[argument]) => argument,
("sum", &[argument]) if plan.expr_type(argument).is_numeric() => argument,
("avg", &[argument])
if plan.expr_type(argument).is_integer()
|| matches!(plan.expr_type(argument), LogicalType::Float | LogicalType::Double) =>
{
argument
}
_ => return None,
};
if !walk::elementwise(plan, written) {
return None;
}
Some(cast(plan, written, &want, span))
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use rudb_common::Stat;
use rudb_common::bounds::{Bound, End, Spread, Test, Zones};
use rudb_common::stat::Provenance;
use rudb_plan::Plan;
use super::project_all;
use crate::estimate::Facts;
#[derive(Debug)]
struct Stub(&'static str, Stat<u64>);
impl Zones for Stub {
fn column(&self, name: &str) -> Option<usize> {
(name == self.0).then_some(0)
}
fn surviving(&self, _tests: &[Test]) -> Option<u64> {
None
}
fn spread(&self, _tests: &[Test]) -> Option<Spread> {
None
}
fn extreme(&self, _column: usize, _end: End) -> Stat<Bound> {
Stat::Unknown
}
fn nulls(&self, _column: usize) -> Stat<u64> {
self.1
}
}
const WATCHED: &str = concat!(
"Aggregate #1 groups=[#0.0::BIGINT, #0.1::INTEGER] aggregates=[count_star()::BIGINT, ",
"sum(#0.2::SMALLINT)::HUGEINT, avg(#0.3::SMALLINT)::DOUBLE]\n",
" Filter (#0.4::VARCHAR <> ''::VARCHAR)::BOOLEAN\n",
" Get memory.main.hits AS hits #0 [w::BIGINT, ip::INTEGER, r::SMALLINT, x::SMALLINT, ",
"p::VARCHAR]\n",
);
fn counted(rows: u64, distinct: u64) -> Facts {
let mut facts = Facts::new();
facts.record("memory", "main", "hits", rows);
facts.record_distinct("memory", "main", "hits", "w", distinct, Provenance::Dictionary);
facts
}
fn projected(text: &str, stats: &Facts, nulls: u64) -> String {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
let zones = Arc::new(Stub("w", Stat::exact(nulls, Provenance::NullCount)));
plan.set_zones(0, zones as Arc<dyn Zones>);
project_all(&mut plan, stats);
plan.validate().unwrap_or_else(|error| panic!("{text} did not stay valid: {error}"));
let once = plan.to_string();
project_all(&mut plan, stats);
assert_eq!(plan.to_string(), once, "a second run moved the plan again");
once
}
#[test]
fn a_key_with_a_value_per_row_makes_the_aggregate_a_projection() {
assert_eq!(
projected(WATCHED, &counted(1_000, 1_000), 0),
concat!(
"Project #1 [#0.0::BIGINT AS column0, #0.1::INTEGER AS column1, 1::BIGINT AS ",
"column2, CAST(#0.2::SMALLINT)::HUGEINT AS column3, ",
"CAST(#0.3::SMALLINT)::DOUBLE AS column4]\n",
" Filter (#0.4::VARCHAR <> ''::VARCHAR)::BOOLEAN\n",
" Get memory.main.hits AS hits #0 [w::BIGINT, ip::INTEGER, r::SMALLINT, ",
"x::SMALLINT, p::VARCHAR]\n",
)
);
}
#[test]
fn a_key_a_projection_passes_on_is_still_a_key() {
let viewed = concat!(
"Aggregate #2 groups=[#1.0::BIGINT, #1.1::INTEGER] aggregates=[count_star()::BIGINT]\n",
" Project #1 [#0.0::BIGINT AS w, \"+\"(#0.1::INTEGER, 1::INTEGER)::INTEGER AS ip]\n",
" Get memory.main.hits AS hits #0 [w::BIGINT, ip::INTEGER, r::SMALLINT, x::SMALLINT, ",
"p::VARCHAR]\n",
);
assert_eq!(
projected(viewed, &counted(1_000, 1_000), 0),
concat!(
"Project #2 [#1.0::BIGINT AS column0, #1.1::INTEGER AS column1, 1::BIGINT AS ",
"column2]\n",
" Project #1 [#0.0::BIGINT AS w, \"+\"(#0.1::INTEGER, 1::INTEGER)::INTEGER AS ip]\n",
" Get memory.main.hits AS hits #0 [w::BIGINT, ip::INTEGER, r::SMALLINT, ",
"x::SMALLINT, p::VARCHAR]\n",
)
);
let computed =
viewed.replace("#0.0::BIGINT AS w", "\"+\"(#0.0::BIGINT, 1::BIGINT)::BIGINT AS w");
assert_eq!(projected(&computed, &counted(1_000, 1_000), 0), computed);
}
#[test]
fn a_key_that_may_repeat_keeps_its_aggregate() {
assert_eq!(projected(WATCHED, &counted(1_000, 999), 0), WATCHED);
assert_eq!(projected(WATCHED, &counted(1_000, 1_000), 1), WATCHED);
let mut rows_only = Facts::new();
rows_only.record("memory", "main", "hits", 1_000);
assert_eq!(projected(WATCHED, &rows_only, 0), WATCHED);
}
#[test]
fn a_key_it_does_not_hold_or_a_call_it_cannot_write_keeps_its_aggregate() {
let other = WATCHED.replace("groups=[#0.0::BIGINT, ", "groups=[");
assert_eq!(projected(&other, &counted(1_000, 1_000), 0), other);
let counts = WATCHED.replace("count_star()::BIGINT", "count(#0.2::SMALLINT)::BIGINT");
assert_eq!(projected(&counts, &counted(1_000, 1_000), 0), counts);
}
const JOINED: &str = concat!(
"Aggregate #5 groups=[#0.0::BIGINT, #0.1::VARCHAR, #2.1::VARCHAR] ",
"aggregates=[sum(#4.1::DECIMAL(38,2))::DECIMAL(38,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 = #4.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT, name::VARCHAR, n::INTEGER]\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",
);
fn joined(text: &str, customers: u64, nations: u64) -> String {
let mut facts = Facts::new();
facts.record("memory", "main", "c", 100);
facts.record_distinct("memory", "main", "c", "k", customers, Provenance::Dictionary);
facts.record("memory", "main", "n", 25);
facts.record_distinct("memory", "main", "n", "k", nations, Provenance::Dictionary);
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
for index in [0, 2] {
let zones = Arc::new(Stub("k", Stat::exact(0, Provenance::NullCount)));
plan.set_zones(index, zones as Arc<dyn Zones>);
}
project_all(&mut plan, &facts);
plan.validate().unwrap_or_else(|error| panic!("{text} did not stay valid: {error}"));
plan.to_string()
}
#[test]
fn a_key_each_join_hands_up_once_makes_the_aggregate_over_them_a_projection() {
let out = joined(JOINED, 100, 25);
assert!(out.starts_with("Project #5 [#0.0::BIGINT AS column0, "), "{out}");
assert!(!out.contains("Aggregate #5"), "{out}");
assert!(out.contains("Aggregate #4"), "{out}");
}
#[test]
fn a_join_that_may_hand_a_row_up_twice_keeps_its_aggregate() {
assert_eq!(joined(JOINED, 99, 25), JOINED);
assert_eq!(joined(JOINED, 100, 24), JOINED);
let wide =
JOINED.replace("groups=[#1.0::BIGINT]", "groups=[#1.0::BIGINT, #1.1::DECIMAL(15,2)]");
assert_eq!(joined(&wide, 100, 25), wide);
let padded = concat!(
"Aggregate #3 groups=[#2.0::INTEGER] aggregates=[count_star()::BIGINT]\n",
" Join LEFT on=[(#0.2::INTEGER = #2.0::INTEGER)::BOOLEAN]\n",
" Get memory.main.c AS c #0 [k::BIGINT, name::VARCHAR, n::INTEGER]\n",
" Get memory.main.n AS n #2 [k::INTEGER, name::VARCHAR]\n",
);
assert_eq!(joined(padded, 100, 25), padded);
}
}