use std::collections::BTreeSet;
use crate::NodeRef;
use crate::node::{BuildSide, JoinKind, Node};
use crate::plan::Plan;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Carried {
tables: BTreeSet<u32>,
}
impl Carried {
#[must_use]
pub fn none() -> Self {
Self::default()
}
#[must_use]
pub fn of(table: u32) -> Self {
Self { tables: BTreeSet::from([table]) }
}
#[must_use]
pub fn has(&self, table: u32) -> bool {
self.tables.contains(&table)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.tables.is_empty()
}
pub fn tables(&self) -> impl Iterator<Item = u32> + '_ {
self.tables.iter().copied()
}
#[must_use]
fn and(mut self, other: &Self) -> Self {
self.tables.extend(other.tables.iter().copied());
self
}
}
#[must_use]
pub fn rids_of(plan: &Plan) -> Vec<Carried> {
let mut known = vec![Carried::none(); 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: &[Carried]) -> Carried {
let below = |node: NodeRef| known[node as usize].clone();
match *plan.node(at) {
Node::Get { index, .. } => Carried::of(index),
Node::TableFetch { index, .. } => Carried::of(index),
Node::Fetch { .. } => Carried::none(),
Node::Filter { input, .. } => below(input),
Node::Project { input, .. } => below(input),
Node::Window { input, .. } => below(input),
Node::Limit { input, .. } | Node::LimitPercent { input, .. } => below(input),
Node::MaterializedCte { body, .. } => below(body),
Node::Sort { .. } | Node::TopN { .. } => Carried::none(),
Node::Distinct { .. } => Carried::none(),
Node::Aggregate { .. } => Carried::none(),
Node::Dummy
| Node::Values { .. }
| Node::TableFunction { .. }
| Node::LateralFunction { .. }
| Node::CteScan { .. } => Carried::none(),
Node::SetOp { .. } => Carried::none(),
Node::CrossProduct { .. } => Carried::none(),
Node::DependentJoin { .. } => Carried::none(),
Node::Join { left, right, kind, build, .. } => {
joined(&below(left), &below(right), kind, build)
}
Node::LinkJoin { child, .. } => below(child),
}
}
fn joined(left: &Carried, right: &Carried, kind: JoinKind, build: BuildSide) -> Carried {
if kind == JoinKind::Positional {
return left.clone().and(right);
}
let probe = match build {
BuildSide::Left => Side::Right,
BuildSide::Right => Side::Left,
};
let padded = match kind {
JoinKind::Left | JoinKind::Single | JoinKind::Mark => Some(Side::Right),
JoinKind::Right => Some(Side::Left),
JoinKind::Full => return Carried::none(),
JoinKind::Inner | JoinKind::Semi | JoinKind::Anti => None,
JoinKind::Positional => unreachable!("answered above"),
};
if padded == Some(probe) {
return Carried::none();
}
if matches!(kind, JoinKind::Semi | JoinKind::Anti) && probe == Side::Right {
return Carried::none();
}
match probe {
Side::Left => left.clone(),
Side::Right => right.clone(),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Side {
Left,
Right,
}
#[cfg(test)]
mod tests {
use rudb_common::{Field, LogicalType, Value};
use super::{Carried, rids_of};
use crate::expr::{ColumnBinding, Expr};
use crate::node::{
Bound, BuildSide, JoinKind, Node, SetOpKind, Share, WindowBound, WindowExclude,
WindowFrame, WindowUnit,
};
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) -> Carried {
rids_of(plan)[plan.root() as usize].clone()
}
fn carried(plan: &Plan) -> Vec<u32> {
root(plan).tables().collect()
}
#[test]
fn a_scan_is_where_a_row_id_comes_from() {
let mut plan = Plan::new();
let node = scan(&mut plan, 7);
plan.set_root(node);
assert_eq!(carried(&plan), vec![7], "the index the scan's columns bind against");
assert!(root(&plan).has(7));
assert!(!root(&plan).has(0), "and no other scan's");
}
#[test]
fn a_filter_and_a_projection_and_a_window_keep_the_rows_they_were_given() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let predicate = column(&mut plan, 0, 0);
let filter = plan.add_node(Node::Filter { input, predicate });
let exprs = plan.add_expr_list(&[predicate]);
let a = plan.intern("a");
let names = plan.add_name_list(&[a]);
let project = plan.add_node(Node::Project { input: filter, index: 1, exprs, names });
let empty = plan.add_expr_list(&[]);
let order = plan.add_sort_keys(&[]);
let window = plan.add_node(Node::Window {
input: project,
index: 2,
partition: empty,
order,
frame: WindowFrame {
unit: WindowUnit::Rows,
start: WindowBound::UnboundedPreceding,
end: WindowBound::CurrentRow,
exclude: WindowExclude::NoOthers,
},
expressions: empty,
});
plan.set_root(window);
assert_eq!(carried(&plan), vec![0]);
}
#[test]
fn a_limit_keeps_them_and_a_limit_percent_does_too() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let limit =
plan.add_node(Node::Limit { input, count: Bound::Rows(10), offset: Bound::All });
let percent = plan.add_node(Node::LimitPercent {
input: limit,
percent: Share::Percent(30.0),
offset: Bound::Rows(0),
});
plan.set_root(percent);
assert_eq!(carried(&plan), vec![0], "a prefix of the rows is still those rows");
}
#[test]
fn a_sort_and_a_top_n_drop_it_until_something_carries_it_through() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let keys = plan.add_sort_keys(&[]);
let sort = plan.add_node(Node::Sort { input, keys });
plan.set_root(sort);
assert!(root(&plan).is_empty());
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let keys = plan.add_sort_keys(&[]);
let top = plan.add_node(Node::TopN { input, keys, count: 10, offset: 0 });
plan.set_root(top);
assert!(root(&plan).is_empty());
}
#[test]
fn an_aggregate_and_a_distinct_produce_rows_that_are_nobody_in_particular() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let empty = plan.add_expr_list(&[]);
let group = column(&mut plan, 0, 0);
let groups = plan.add_expr_list(&[group]);
let node = plan.add_node(Node::Aggregate { input, index: 1, groups, aggregates: empty });
plan.set_root(node);
assert!(root(&plan).is_empty(), "a group is not a row of anything");
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let on = plan.add_expr_list(&[]);
let node = plan.add_node(Node::Distinct { input, on });
plan.set_root(node);
assert!(root(&plan).is_empty(), "which row of a duplicate group survived is not defined");
}
#[test]
fn rows_that_were_never_a_table_carry_nothing() {
let mut plan = Plan::new();
let node = plan.add_node(Node::Dummy);
plan.set_root(node);
assert!(root(&plan).is_empty());
let mut plan = Plan::new();
let columns = plan.add_fields(&[Field::new("a", LogicalType::Integer)]);
let one = plan.add_constant(Value::Integer(1));
let row = plan.add_expr_list(&[one]);
let rows = plan.add_rows(&[row]);
let node = plan.add_node(Node::Values { index: 0, columns, rows });
plan.set_root(node);
assert!(root(&plan).is_empty(), "a literal row is not a row of a table");
let mut plan = Plan::new();
let columns = plan.add_fields(&[Field::new("a", LogicalType::Integer)]);
let name = plan.intern("range");
let empty = plan.add_expr_list(&[]);
let node = plan.add_node(Node::TableFunction {
index: 0,
function: name,
args: empty,
options: empty,
settings: empty,
columns,
});
plan.set_root(node);
assert!(root(&plan).is_empty());
let mut plan = Plan::new();
let columns = plan.add_fields(&[Field::new("a", LogicalType::Integer)]);
let name = plan.intern("x");
let node = plan.add_node(Node::CteScan { index: 0, cte: 0, name, columns });
plan.set_root(node);
assert!(root(&plan).is_empty(), "what it reads was computed rather than stored");
}
#[test]
fn a_lateral_function_produces_its_own_rows_and_not_its_inputs() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let columns = plan.add_fields(&[Field::new("a", LogicalType::Integer)]);
let name = plan.intern("unnest");
let empty = plan.add_expr_list(&[]);
let node = plan.add_node(Node::LateralFunction {
input,
index: 1,
function: name,
args: empty,
options: empty,
settings: empty,
columns,
});
plan.set_root(node);
assert!(root(&plan).is_empty(), "one input row becomes any number of output rows");
}
#[test]
fn a_table_fetch_is_a_row_id_and_a_file_fetch_is_not() {
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let row = column(&mut plan, 0, 0);
let columns = plan.add_fields(&[Field::new("a", LogicalType::Integer)]);
let name = plan.intern("t");
let node = plan.add_node(Node::TableFetch {
input,
index: 3,
catalog: name,
schema: name,
table: name,
columns,
row,
});
plan.set_root(node);
assert_eq!(carried(&plan), vec![3], "a row read back by its ordinal is that row");
let mut plan = Plan::new();
let input = scan(&mut plan, 0);
let row = column(&mut plan, 0, 0);
let columns = plan.add_fields(&[Field::new("a", LogicalType::Integer)]);
let args = plan.add_expr_list(&[]);
let node = plan.add_node(Node::Fetch { input, index: 3, args, columns, row });
plan.set_root(node);
assert!(root(&plan).is_empty(), "an ordinal in a file is not a row id of a table");
}
#[test]
fn a_materialisation_carries_what_its_body_carries_and_a_set_operation_carries_nothing() {
let mut plan = Plan::new();
let definition = scan(&mut plan, 0);
let body = scan(&mut plan, 1);
let name = plan.intern("x");
let columns = plan.add_fields(&[Field::new("a", LogicalType::Integer)]);
let node = plan.add_node(Node::MaterializedCte { definition, body, name, cte: 0, columns });
plan.set_root(node);
assert_eq!(carried(&plan), vec![1], "the rows are the body's, not the definition's");
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: SetOpKind::Union, all: true, index: 2 });
plan.set_root(node);
assert!(root(&plan).is_empty(), "two tables' rows under one schema are neither table's");
}
#[test]
fn a_cross_product_and_a_dependent_join_carry_nothing() {
let mut plan = Plan::new();
let left = scan(&mut plan, 0);
let right = scan(&mut plan, 1);
let node = plan.add_node(Node::CrossProduct { left, right });
plan.set_root(node);
assert!(root(&plan).is_empty());
let mut plan = Plan::new();
let left = scan(&mut plan, 0);
let right = scan(&mut plan, 1);
let conditions = plan.add_expr_list(&[]);
let node =
plan.add_node(Node::DependentJoin { left, right, kind: JoinKind::Inner, conditions });
plan.set_root(node);
assert!(root(&plan).is_empty(), "it has to be unnested before it can run at all");
}
fn join(kind: JoinKind, build: BuildSide) -> Plan {
let mut plan = Plan::new();
let left = scan(&mut plan, 0);
let right = scan(&mut plan, 1);
let conditions = plan.add_expr_list(&[]);
let node = plan.add_node(Node::Join { left, right, kind, conditions, build });
plan.set_root(node);
plan
}
#[test]
fn a_join_carries_the_side_it_streams_and_not_the_side_it_gathers() {
assert_eq!(carried(&join(JoinKind::Inner, BuildSide::Right)), vec![0]);
assert_eq!(carried(&join(JoinKind::Inner, BuildSide::Left)), vec![1]);
}
#[test]
fn an_outer_join_carries_nothing_for_a_side_it_pads() {
assert!(root(&join(JoinKind::Left, BuildSide::Left)).is_empty());
assert_eq!(
carried(&join(JoinKind::Left, BuildSide::Right)),
vec![0],
"the left side of a left join is never padded, and here it is also the probe"
);
assert!(root(&join(JoinKind::Right, BuildSide::Right)).is_empty());
assert_eq!(carried(&join(JoinKind::Right, BuildSide::Left)), vec![1]);
assert!(root(&join(JoinKind::Full, BuildSide::Right)).is_empty());
assert!(root(&join(JoinKind::Full, BuildSide::Left)).is_empty());
}
#[test]
fn a_semi_join_carries_its_left_side_only_where_the_left_side_probes() {
assert_eq!(carried(&join(JoinKind::Semi, BuildSide::Right)), vec![0]);
assert!(root(&join(JoinKind::Semi, BuildSide::Left)).is_empty());
assert_eq!(carried(&join(JoinKind::Anti, BuildSide::Right)), vec![0]);
assert!(root(&join(JoinKind::Anti, BuildSide::Left)).is_empty());
}
#[test]
fn a_single_and_a_mark_join_pad_the_right_side() {
assert_eq!(carried(&join(JoinKind::Single, BuildSide::Right)), vec![0]);
assert!(root(&join(JoinKind::Single, BuildSide::Left)).is_empty());
assert_eq!(carried(&join(JoinKind::Mark, BuildSide::Right)), vec![0]);
assert!(root(&join(JoinKind::Mark, BuildSide::Left)).is_empty());
}
#[test]
fn a_positional_join_carries_both_sides() {
assert_eq!(carried(&join(JoinKind::Positional, BuildSide::Right)), vec![0, 1]);
assert_eq!(carried(&join(JoinKind::Positional, BuildSide::Left)), vec![0, 1]);
}
#[test]
fn a_link_join_carries_its_child_and_never_its_parent() {
for kind in [JoinKind::Inner, JoinKind::Left, JoinKind::Semi, JoinKind::Anti] {
let mut plan = Plan::new();
let child = scan(&mut plan, 4);
let parent = scan(&mut plan, 5);
let conditions = plan.add_expr_list(&[]);
let rid = column(&mut plan, 4, 0);
let node = plan.add_node(Node::LinkJoin { child, parent, kind, conditions, rid });
plan.set_root(node);
assert_eq!(carried(&plan), vec![4], "{kind:?} keeps the child and only the child");
assert!(!root(&plan).has(5), "{kind:?} gathered the parent");
}
}
#[test]
fn a_row_id_survives_a_stack_of_operators_that_all_preserve_it() {
let mut plan = Plan::new();
let orders = scan(&mut plan, 0);
let predicate = column(&mut plan, 0, 0);
let filter = plan.add_node(Node::Filter { input: orders, predicate });
let exprs = plan.add_expr_list(&[predicate]);
let a = plan.intern("a");
let names = plan.add_name_list(&[a]);
let project = plan.add_node(Node::Project { input: filter, index: 1, exprs, names });
let lineitem = scan(&mut plan, 2);
let conditions = plan.add_expr_list(&[]);
let node = plan.add_node(Node::Join {
left: lineitem,
right: project,
kind: JoinKind::Inner,
conditions,
build: BuildSide::Right,
});
plan.set_root(node);
assert_eq!(carried(&plan), vec![2], "the child side streams, so its row ids are there");
let all = rids_of(&plan);
assert_eq!(all[project as usize].tables().collect::<Vec<_>>(), vec![0]);
assert_eq!(all[filter as usize].tables().collect::<Vec<_>>(), vec![0]);
}
}