use std::collections::BTreeSet;
use crate::expr::{ColumnBinding, CompareOp, ConjunctionOp, Expr};
use crate::node::{Bound, JoinKind, Node};
use crate::plan::Plan;
use crate::{ExprRef, NodeRef, Slice};
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Keys {
sets: Vec<Vec<ColumnBinding>>,
row: bool,
constants: BTreeSet<ColumnBinding>,
}
impl Keys {
#[must_use]
pub fn unknown() -> Self {
Self::default()
}
#[must_use]
pub fn single() -> Self {
Self { sets: vec![Vec::new()], row: true, constants: BTreeSet::new() }
}
#[must_use]
pub fn of(columns: impl IntoIterator<Item = ColumnBinding>) -> Self {
let mut keys = Self::default();
keys.add(columns.into_iter().collect());
keys
}
#[must_use]
pub fn whole_row() -> Self {
Self { sets: Vec::new(), row: true, constants: BTreeSet::new() }
}
#[must_use]
pub fn sets(&self) -> &[Vec<ColumnBinding>] {
&self.sets
}
#[must_use]
pub const fn row(&self) -> bool {
self.row
}
pub fn constants(&self) -> impl Iterator<Item = ColumnBinding> + '_ {
self.constants.iter().copied()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.sets.is_empty() && !self.row && self.constants.is_empty()
}
#[must_use]
pub fn at_most_one_row(&self) -> bool {
self.sets.first().is_some_and(Vec::is_empty)
}
#[must_use]
pub fn covers(&self, columns: &[ColumnBinding]) -> bool {
let held: BTreeSet<ColumnBinding> =
columns.iter().copied().filter(|column| !self.constants.contains(column)).collect();
self.sets.iter().any(|set| set.iter().all(|column| held.contains(column)))
}
fn add(&mut self, mut columns: Vec<ColumnBinding>) {
columns.sort_unstable();
columns.dedup();
columns.retain(|column| !self.constants.contains(column));
if self.sets.iter().any(|set| set.iter().all(|column| columns.contains(column))) {
return;
}
self.sets.retain(|set| !columns.iter().all(|column| set.contains(column)));
self.sets.push(columns);
self.sets.sort_by(|a, b| a.len().cmp(&b.len()).then_with(|| a.cmp(b)));
}
fn fix(&mut self, column: ColumnBinding) {
if !self.constants.insert(column) {
return;
}
let sets = std::mem::take(&mut self.sets);
for set in sets {
self.add(set);
}
}
fn product(&self, other: &Self) -> Self {
let mut keys = Self::default();
for column in self.constants.iter().chain(other.constants.iter()) {
keys.constants.insert(*column);
}
for left in &self.sets {
for right in &other.sets {
let mut both = left.clone();
both.extend(right.iter().copied());
keys.add(both);
}
}
keys
}
}
#[must_use]
pub fn keys_of(plan: &Plan) -> Vec<Keys> {
let mut known = vec![Keys::unknown(); plan.node_count()];
let mut order = Vec::with_capacity(plan.node_count());
postorder(plan, plan.root(), &mut order);
for node in order {
known[node as usize] = compute(plan, node, &known);
}
known
}
fn postorder(plan: &Plan, at: NodeRef, out: &mut Vec<NodeRef>) {
for child in plan.node(at).children().into_iter().flatten() {
postorder(plan, child, out);
}
out.push(at);
}
fn compute(plan: &Plan, at: NodeRef, known: &[Keys]) -> Keys {
let below = |node: NodeRef| known[node as usize].clone();
match *plan.node(at) {
Node::Dummy => Keys::single(),
Node::Aggregate { input, index, groups, .. } => {
let source = below(input);
let mut keys = Keys::default();
for (position, expr) in plan.expr_list(groups).iter().enumerate() {
let Expr::Column(binding) = *plan.expr(*expr) else { continue };
if source.constants.contains(&binding) {
keys.constants.insert(ColumnBinding::new(index, at_most_u32(position)));
}
}
keys.add(
(0..len(plan, groups))
.map(|position| ColumnBinding::new(index, position))
.collect(),
);
keys
}
Node::Distinct { input, on } => {
let mut keys = below(input);
if on.len == 0 {
keys.row = true;
return keys;
}
let exprs = plan.expr_list(on);
let mut columns = Vec::with_capacity(exprs.len());
for expr in exprs {
let Expr::Column(binding) = *plan.expr(*expr) else { return keys };
columns.push(binding);
}
keys.add(columns);
keys
}
Node::SetOp { all, .. } => {
if all {
Keys::unknown()
} else {
Keys::whole_row()
}
}
Node::Filter { input, predicate } => {
let mut keys = below(input);
for column in fixed(plan, predicate) {
keys.fix(column);
}
keys
}
Node::Sort { input, .. } | Node::Window { input, .. } => below(input),
Node::Limit { input, count, .. } => match count {
Bound::Rows(0 | 1) => Keys::single(),
Bound::All | Bound::Rows(_) | Bound::Read(_) => below(input),
},
Node::LimitPercent { input, .. } => below(input),
Node::TopN { input, count, .. } => {
if count <= 1 {
Keys::single()
} else {
below(input)
}
}
Node::Project { input, index, exprs, .. } => {
let mut moved = Vec::new();
for (position, expr) in plan.expr_list(exprs).iter().enumerate() {
if let Expr::Column(binding) = *plan.expr(*expr) {
moved.push((binding, ColumnBinding::new(index, at_most_u32(position))));
}
}
let source = below(input);
let mut keys = Keys::default();
for (from, to) in &moved {
if source.constants.contains(from) {
keys.constants.insert(*to);
}
}
for set in &source.sets {
let mut mapped = Vec::with_capacity(set.len());
for column in set {
let Some((_, to)) = moved.iter().find(|(from, _)| from == column) else {
mapped.clear();
break;
};
mapped.push(*to);
}
if mapped.len() == set.len() {
keys.add(mapped);
}
}
keys
}
Node::Join { left, right, kind, conditions, build: _ } => match kind {
JoinKind::Inner | JoinKind::Positional => {
let (left, right) = (below(left), below(right));
let mut keys = left.product(&right);
for (from, onto) in [(&left, &right), (&right, &left)] {
if onto.covers(&equated(plan, conditions)) {
for set in &from.sets {
keys.add(set.clone());
}
}
}
keys
}
JoinKind::Semi | JoinKind::Anti | JoinKind::Mark | JoinKind::Single => below(left),
JoinKind::Left | JoinKind::Right | JoinKind::Full => Keys::unknown(),
},
Node::CrossProduct { left, right } => below(left).product(&below(right)),
Node::MaterializedCte { body, .. } => below(body),
Node::Get { .. }
| Node::Values { .. }
| Node::TableFunction { .. }
| Node::LateralFunction { .. }
| Node::Fetch { .. }
| Node::TableFetch { .. }
| Node::CteScan { .. }
| Node::DependentJoin { .. } => Keys::unknown(),
}
}
fn equated(plan: &Plan, conditions: Slice) -> Vec<ColumnBinding> {
let mut found = Vec::new();
for condition in plan.expr_list(conditions) {
let Expr::Compare { op: CompareOp::Equal, left, right } = *plan.expr(*condition) else {
continue;
};
for side in [left, right] {
if let Expr::Column(binding) = *plan.expr(side) {
found.push(binding);
}
}
}
found
}
fn fixed(plan: &Plan, predicate: ExprRef) -> Vec<ColumnBinding> {
let mut found = Vec::new();
collect_fixed(plan, predicate, &mut found);
found
}
fn collect_fixed(plan: &Plan, predicate: ExprRef, found: &mut Vec<ColumnBinding>) {
match *plan.expr(predicate) {
Expr::Conjunction { op: ConjunctionOp::And, children } => {
for child in plan.expr_list(children) {
collect_fixed(plan, *child, found);
}
}
Expr::Compare { op: CompareOp::Equal | CompareOp::NotDistinctFrom, left, right } => {
for (column, other) in [(left, right), (right, left)] {
if let (Expr::Column(binding), Expr::Constant(_)) =
(plan.expr(column), plan.expr(other))
{
found.push(*binding);
}
}
}
_ => {}
}
}
fn len(plan: &Plan, slice: Slice) -> u32 {
at_most_u32(plan.expr_list(slice).len())
}
fn at_most_u32(position: usize) -> u32 {
u32::try_from(position).unwrap_or(u32::MAX)
}
#[cfg(test)]
mod tests {
use rudb_common::{Field, LogicalType, Value};
use super::{Keys, keys_of};
use crate::expr::{ColumnBinding, CompareOp, Expr};
use crate::node::{Bound, JoinKind, Node};
use crate::plan::Plan;
fn scan(plan: &mut Plan, index: u32) -> u32 {
let columns = plan.add_fields(&[
Field::new("a", LogicalType::Integer),
Field::new("b", LogicalType::Integer),
]);
let name = plan.intern("t");
plan.add_node(Node::Get {
catalog: name,
schema: name,
table: name,
alias: name,
index,
columns,
})
}
fn column(plan: &mut Plan, table: u32, position: u32) -> u32 {
plan.add_expr(Expr::Column(ColumnBinding::new(table, position)), LogicalType::Integer)
}
fn root(plan: &Plan) -> Keys {
keys_of(plan)[plan.root() as usize].clone()
}
#[test]
fn a_group_by_is_keyed_by_what_it_grouped_on() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let group = column(&mut plan, 0, 0);
let groups = plan.add_expr_list(&[group]);
let empty = plan.add_expr_list(&[]);
let node = plan.add_node(Node::Aggregate { input, index: 1, groups, aggregates: empty });
plan.set_root(node);
assert_eq!(root(&plan).sets(), [vec![ColumnBinding::new(1, 0)]]);
}
#[test]
fn an_ungrouped_aggregate_produces_one_row_and_says_so_with_the_empty_key() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let empty = plan.add_expr_list(&[]);
let node =
plan.add_node(Node::Aggregate { input, index: 1, groups: empty, aggregates: empty });
plan.set_root(node);
assert!(root(&plan).at_most_one_row());
}
#[test]
fn a_scan_claims_nothing_because_nothing_stores_a_primary_key() {
let mut plan = Plan::new();
let node = scan(&mut plan, 0);
plan.set_root(node);
assert!(root(&plan).is_empty());
}
#[test]
fn a_column_held_equal_to_a_literal_comes_out_of_the_key_it_was_in() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let a = column(&mut plan, 0, 0);
let b = column(&mut plan, 0, 1);
let groups = plan.add_expr_list(&[a, b]);
let empty = plan.add_expr_list(&[]);
let agg = plan.add_node(Node::Aggregate { input, index: 1, groups, aggregates: empty });
let left = column(&mut plan, 1, 0);
let one = plan.add_constant(Value::Integer(1));
let predicate = plan.add_expr(
Expr::Compare { op: CompareOp::Equal, left, right: one },
LogicalType::Boolean,
);
let filter = plan.add_node(Node::Filter { input: agg, predicate });
plan.set_root(filter);
let all = keys_of(&plan);
assert_eq!(
all[agg as usize].sets(),
[vec![ColumnBinding::new(1, 0), ColumnBinding::new(1, 1)]]
);
let keys = all[filter as usize].clone();
assert_eq!(keys.constants().collect::<Vec<_>>(), [ColumnBinding::new(1, 0)]);
assert_eq!(keys.sets(), [vec![ColumnBinding::new(1, 1)]]);
assert!(keys.covers(&[ColumnBinding::new(1, 0), ColumnBinding::new(1, 1)]));
assert!(!keys.at_most_one_row());
}
#[test]
fn a_group_by_on_a_column_the_filter_under_it_fixed_produces_one_row() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let left = column(&mut plan, 0, 0);
let one = plan.add_constant(Value::Integer(1));
let predicate = plan.add_expr(
Expr::Compare { op: CompareOp::Equal, left, right: one },
LogicalType::Boolean,
);
let filter = plan.add_node(Node::Filter { input, predicate });
let group = column(&mut plan, 0, 0);
let groups = plan.add_expr_list(&[group]);
let empty = plan.add_expr_list(&[]);
let node =
plan.add_node(Node::Aggregate { input: filter, index: 1, groups, aggregates: empty });
plan.set_root(node);
assert!(root(&plan).at_most_one_row());
}
#[test]
fn a_predicate_that_fixes_every_column_of_a_key_leaves_at_most_one_row() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let a = column(&mut plan, 0, 0);
let groups = plan.add_expr_list(&[a]);
let empty = plan.add_expr_list(&[]);
let agg = plan.add_node(Node::Aggregate { input, index: 1, groups, aggregates: empty });
let left = column(&mut plan, 1, 0);
let one = plan.add_constant(Value::Integer(1));
let predicate = plan.add_expr(
Expr::Compare { op: CompareOp::Equal, left, right: one },
LogicalType::Boolean,
);
let filter = plan.add_node(Node::Filter { input: agg, predicate });
plan.set_root(filter);
assert!(root(&plan).at_most_one_row());
}
#[test]
fn an_or_fixes_nothing_because_a_column_is_pinned_by_one_arm_and_free_in_the_other() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let a = column(&mut plan, 0, 0);
let groups = plan.add_expr_list(&[a]);
let empty = plan.add_expr_list(&[]);
let agg = plan.add_node(Node::Aggregate { input, index: 1, groups, aggregates: empty });
let mut arms = Vec::new();
for value in [1i32, 2] {
let left = column(&mut plan, 1, 0);
let literal = plan.add_constant(Value::Integer(value));
arms.push(plan.add_expr(
Expr::Compare { op: CompareOp::Equal, left, right: literal },
LogicalType::Boolean,
));
}
let children = plan.add_expr_list(&arms);
let predicate = plan.add_expr(
Expr::Conjunction { op: crate::expr::ConjunctionOp::Or, children },
LogicalType::Boolean,
);
let filter = plan.add_node(Node::Filter { input: agg, predicate });
plan.set_root(filter);
let keys = root(&plan);
assert_eq!(keys.constants().count(), 0);
assert!(!keys.at_most_one_row());
}
#[test]
fn a_projection_carries_a_key_through_and_drops_one_it_did_not_pass_on() {
for (kept, expected) in [(0u32, 1usize), (1, 0)] {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let group = column(&mut plan, 0, 0);
let groups = plan.add_expr_list(&[group]);
let empty = plan.add_expr_list(&[]);
let agg = plan.add_node(Node::Aggregate { input, index: 1, groups, aggregates: empty });
let passed = column(&mut plan, 1, kept);
let exprs = plan.add_expr_list(&[passed]);
let name = plan.intern("x");
let names = plan.add_name_list(&[name]);
let node = plan.add_node(Node::Project { input: agg, index: 2, exprs, names });
plan.set_root(node);
assert_eq!(root(&plan).sets().len(), expected, "keeping column {kept}");
}
}
#[test]
fn a_join_onto_a_key_keeps_the_other_sides_keys_on_their_own() {
let mut plan = Plan::new();
let left_scan = scan(&mut plan, 0);
let left_group = column(&mut plan, 0, 0);
let left_groups = plan.add_expr_list(&[left_group]);
let empty = plan.add_expr_list(&[]);
let left = plan.add_node(Node::Aggregate {
input: left_scan,
index: 1,
groups: left_groups,
aggregates: empty,
});
let right_scan = scan(&mut plan, 2);
let right_group = column(&mut plan, 2, 0);
let right_groups = plan.add_expr_list(&[right_group]);
let right = plan.add_node(Node::Aggregate {
input: right_scan,
index: 3,
groups: right_groups,
aggregates: empty,
});
let on_left = column(&mut plan, 1, 0);
let on_right = column(&mut plan, 3, 0);
let condition = plan.add_expr(
Expr::Compare { op: CompareOp::Equal, left: on_left, right: on_right },
LogicalType::Boolean,
);
let conditions = plan.add_expr_list(&[condition]);
let node = plan.add_node(Node::Join {
left,
right,
kind: JoinKind::Inner,
conditions,
build: crate::node::BuildSide::Right,
});
plan.set_root(node);
let keys = root(&plan);
assert!(keys.covers(&[ColumnBinding::new(1, 0)]), "{keys:?}");
assert!(keys.covers(&[ColumnBinding::new(3, 0)]), "{keys:?}");
}
#[test]
fn an_outer_join_keeps_nothing_because_two_padded_rows_agree_where_the_key_was() {
let mut plan = Plan::new();
let left_scan = scan(&mut plan, 0);
let group = column(&mut plan, 0, 0);
let groups = plan.add_expr_list(&[group]);
let empty = plan.add_expr_list(&[]);
let left = plan.add_node(Node::Aggregate {
input: left_scan,
index: 1,
groups,
aggregates: empty,
});
let right = scan(&mut plan, 2);
let node = plan.add_node(Node::Join {
left,
right,
kind: JoinKind::Left,
conditions: empty,
build: crate::node::BuildSide::Right,
});
plan.set_root(node);
assert!(root(&plan).is_empty());
}
#[test]
fn a_distinct_keys_the_whole_row_and_a_union_all_keys_nothing() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let empty = plan.add_expr_list(&[]);
let node = plan.add_node(Node::Distinct { input, on: empty });
plan.set_root(node);
assert!(root(&plan).row());
let mut plan = Plan::new();
let left = scan(&mut plan, 0);
let right = scan(&mut plan, 1);
let node = plan.add_node(Node::SetOp {
left,
right,
kind: crate::node::SetOpKind::Union,
all: true,
index: 2,
});
plan.set_root(node);
assert!(root(&plan).is_empty());
}
#[test]
fn a_limit_of_one_row_is_one_row_whatever_was_underneath_it() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let node =
plan.add_node(Node::Limit { input, count: Bound::Rows(1), offset: Bound::Rows(0) });
plan.set_root(node);
assert!(root(&plan).at_most_one_row());
}
}