use rudb_common::{Class, LogicalType, Result, Stat, Value};
use rudb_plan::{ColumnBinding, Expr, ExprRef, Node, NodeRef, Plan};
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(())
}
}
pub fn project_all(plan: &mut Plan, stats: &Facts) {
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);
}
}
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(binding) = *plan.expr(key) else { return false };
let Some(scan) = filtered_scan(plan, input, binding) else { return false };
let Node::Get { catalog, schema, table, columns, .. } = *plan.node(scan) 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, input, binding);
never_null && rows.is_some() && rows == values
}
fn filtered_scan(plan: &Plan, at: NodeRef, binding: ColumnBinding) -> Option<NodeRef> {
match *plan.node(at) {
Node::Get { index, .. } if index == binding.table => Some(at),
Node::Filter { input, .. } => filtered_scan(plan, input, binding),
_ => None,
}
}
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(Stat<u64>);
impl Zones for Stub {
fn column(&self, name: &str) -> Option<usize> {
(name == "w").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.0
}
}
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(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_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);
}
}