use rudb_common::Result;
use rudb_common::rules::Rule;
use rudb_plan::{ColumnBinding, Expr, Node, NodeRef, Plan};
use crate::link::Linked;
use crate::pass::{Context, Pass, top_down};
#[derive(Debug)]
pub struct AggregateCluster;
impl Pass for AggregateCluster {
fn name(&self) -> &'static str {
"aggregate_cluster"
}
fn run(&self, plan: &mut Plan, context: &Context) -> Result<()> {
if context.allows(Rule::ClosedGroups) {
cluster(plan, context.links());
}
Ok(())
}
}
fn cluster(plan: &mut Plan, links: &[Linked]) {
let mut found = Vec::new();
for node in top_down(plan) {
let Node::Aggregate { input, index, groups, .. } = *plan.node(node) else {
continue;
};
let &[key] = plan.expr_list(groups) else { continue };
let &Expr::Column(binding) = plan.expr(key) else { continue };
match sorted(plan, input, binding, links, 16) {
Some(Run::Ascending) => found.push((index, false)),
Some(Run::Grouped) => found.push((index, true)),
None => {}
}
}
for (index, grouped) in found {
if grouped {
plan.cluster_grouped(index);
} else {
plan.cluster(index);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Run {
Ascending,
Grouped,
}
fn sorted(
plan: &Plan,
node: NodeRef,
binding: ColumnBinding,
links: &[Linked],
depth: u32,
) -> Option<Run> {
let depth = depth.checked_sub(1)?;
match *plan.node(node) {
Node::Filter { input, .. } => sorted(plan, input, binding, links, depth),
Node::Project { input, index, exprs, .. } if index == binding.table => {
let &carried = plan.expr_list(exprs).get(binding.column as usize)?;
let &Expr::Column(carried) = plan.expr(carried) else { return None };
sorted(plan, input, carried, links, depth)
}
Node::Get { index, table, columns, .. } if index == binding.table => {
let field = plan.field_list(columns).get(binding.column as usize)?;
if plan.ascending(index, &field.name) {
return Some(Run::Ascending);
}
let table = plan.string(table);
links
.iter()
.any(|link| {
link.groups_child()
&& link.child.eq_ignore_ascii_case(table)
&& link.child_column.eq_ignore_ascii_case(&field.name)
})
.then_some(Run::Grouped)
}
_ => None,
}
}
#[cfg(test)]
mod tests {
use rudb_plan::Plan;
use super::{AggregateCluster, Rule};
use crate::pass::{Context, Pass};
const SCAN: &str = "Get memory.main.t AS t #0 [a::INTEGER, b::INTEGER]";
fn plan(text: &str) -> Plan {
let mut plan = Plan::parse(text).expect("a plan that parses");
plan.mark_ascending(0, "a");
plan
}
fn run(plan: &mut Plan, context: &Context) {
AggregateCluster.run(plan, context).expect("a pass that cannot fail");
}
#[test]
fn a_key_the_table_is_sorted_on_is_clustered() {
let mut plan =
plan(&format!("Aggregate #1 groups=[#0.0::INTEGER] aggregates=[]\n {SCAN}\n"));
run(&mut plan, &Context::new());
assert!(plan.clustered(1));
}
#[test]
fn a_key_the_table_is_not_sorted_on_is_left_alone() {
let mut plan =
plan(&format!("Aggregate #1 groups=[#0.1::INTEGER] aggregates=[]\n {SCAN}\n"));
run(&mut plan, &Context::new());
assert!(!plan.clustered(1));
}
#[test]
fn two_keys_are_left_alone() {
let mut plan = plan(&format!(
"Aggregate #1 groups=[#0.0::INTEGER, #0.1::INTEGER] aggregates=[]\n {SCAN}\n"
));
run(&mut plan, &Context::new());
assert!(!plan.clustered(1));
}
#[test]
fn a_filter_keeps_the_order() {
let mut plan = plan(&format!(
"Aggregate #1 groups=[#0.0::INTEGER] aggregates=[]\n Filter (#0.1::INTEGER > 3::INTEGER)::BOOLEAN\n {SCAN}\n"
));
run(&mut plan, &Context::new());
assert!(plan.clustered(1));
}
fn linked(link: crate::link::Linked) -> Context {
let mut context = Context::new();
context.relate(std::sync::Arc::new(vec![link]));
context
}
#[test]
fn a_key_a_monotone_total_link_proves_grouped_is_clustered_as_grouped() {
let text = format!("Aggregate #1 groups=[#0.1::INTEGER] aggregates=[]\n {SCAN}\n");
let proof = crate::link::Linked::verified("t", "b", "p", "k").monotone();
let mut marked = plan(&text);
run(&mut marked, &linked(proof.clone()));
assert!(marked.clustered(1));
assert!(marked.grouped(1), "grouped, so the operator does not expect it to go up");
let mut ascending =
plan(&format!("Aggregate #1 groups=[#0.0::INTEGER] aggregates=[]\n {SCAN}\n"));
run(&mut ascending, &linked(proof.clone()));
assert!(ascending.clustered(1) && !ascending.grouped(1));
for weaker in [
crate::link::Linked::verified("t", "b", "p", "k"),
crate::link::Linked::built("t", "b", "p", "k").monotone(),
proof.clone().and("a", "j"),
crate::link::Linked::verified("t", "a", "p", "k").monotone(),
] {
let mut marked = plan(&text);
run(&mut marked, &linked(weaker.clone()));
assert!(!marked.clustered(1), "{weaker:?}");
}
}
#[test]
fn the_rule_s_own_setting_turns_it_off() {
let mut plan =
plan(&format!("Aggregate #1 groups=[#0.0::INTEGER] aggregates=[]\n {SCAN}\n"));
let mut context = Context::new();
let mut rules = rudb_common::rules::Rules::default();
rules.set(Rule::ClosedGroups, false);
context.govern(rules);
run(&mut plan, &context);
assert!(!plan.clustered(1));
}
}