use std::collections::HashMap;
use rudb_common::{Field, LogicalType, Value};
use rudb_plan::{
Arm, BuildSide, ColumnBinding, CompareOp, ConjunctionOp, Expr, ExprRef, JoinKind, Node,
NodeRef, Plan, Slice, SortKey, StrRef, WindowBound, WindowExclude, WindowFrame, WindowUnit,
};
use crate::tables::{TableSet, produced};
use crate::walk;
struct Key {
binding: ColumnBinding,
expr: ExprRef,
}
struct Pushed {
node: NodeRef,
keys: Vec<ColumnBinding>,
moved: HashMap<ColumnBinding, ColumnBinding>,
}
pub(crate) fn lower(
plan: &mut Plan,
left: NodeRef,
right: NodeRef,
kind: JoinKind,
conditions: Slice,
) -> Option<NodeRef> {
let outer = produced(plan, left);
let keys = read_keys(plan, right, &outer);
if keys.is_empty() {
return None;
}
let group_exprs: Vec<ExprRef> = keys.iter().map(|key| key.expr).collect();
let groups = plan.add_expr_list(&group_exprs);
let aggregates = plan.add_expr_list(&[]);
let index = walk::fresh_index(plan);
let domain = plan.add_node(Node::Aggregate { input: left, index, groups, aggregates });
let pushed = push(plan, right, domain, index, &keys, &outer)?;
let held = plan.expr_list(conditions).to_vec();
let mut all: Vec<ExprRef> =
held.into_iter().map(|condition| remap(plan, condition, &pushed.moved)).collect();
for (key, &carried) in keys.iter().zip(&pushed.keys) {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
let here = plan.add_expr_at(Expr::Column(key.binding), ty.clone(), span);
let there = plan.add_expr_at(Expr::Column(carried), ty, span);
all.push(plan.add_expr_at(
Expr::Compare { op: CompareOp::NotDistinctFrom, left: here, right: there },
LogicalType::Boolean,
span,
));
}
let conditions = plan.add_expr_list(&all);
Some(plan.add_node(Node::Join {
left,
right: pushed.node,
kind,
conditions,
build: BuildSide::default(),
}))
}
fn read_keys(plan: &Plan, at: NodeRef, outer: &TableSet) -> Vec<Key> {
let mut keys: Vec<Key> = Vec::new();
collect_keys(plan, at, outer, &mut keys);
keys
}
fn collect_keys(plan: &Plan, at: NodeRef, outer: &TableSet, keys: &mut Vec<Key>) {
walk::node_columns(plan, at, &mut |expr, binding| {
if outer.contains(binding.table) && !keys.iter().any(|key| key.binding == binding) {
keys.push(Key { binding, expr });
}
});
for child in plan.node(at).children().into_iter().flatten() {
collect_keys(plan, child, outer, keys);
}
}
fn correlated(plan: &Plan, at: NodeRef, outer: &TableSet) -> bool {
let mut yes = false;
walk::node_columns(plan, at, &mut |_, binding| yes |= outer.contains(binding.table));
yes || plan
.node(at)
.children()
.into_iter()
.flatten()
.any(|child| correlated(plan, child, outer))
}
fn push(
plan: &mut Plan,
at: NodeRef,
domain: NodeRef,
index: u32,
keys: &[Key],
outer: &TableSet,
) -> Option<Pushed> {
if !correlated(plan, at, outer) {
let node = plan.add_node(Node::CrossProduct { left: domain, right: at });
let carried = (0..keys.len())
.map(|position| ColumnBinding::new(index, u32::try_from(position).expect("key count")))
.collect();
return Some(Pushed { node, keys: carried, moved: HashMap::new() });
}
match *plan.node(at) {
Node::Filter { input, predicate } => {
let below = push(plan, input, domain, index, keys, outer)?;
let map = mapping(keys, &below);
let predicate = remap(plan, predicate, &map);
let node = plan.add_node(Node::Filter { input: below.node, predicate });
Some(Pushed { node, keys: below.keys, moved: below.moved })
}
Node::Project { input, index: at_index, exprs, names } => {
let below = push(plan, input, domain, index, keys, outer)?;
let map = mapping(keys, &below);
let held = plan.expr_list(exprs).to_vec();
let mut projected: Vec<ExprRef> =
held.into_iter().map(|expr| remap(plan, expr, &map)).collect();
let mut projected_names = plan.name_list(names).to_vec();
let width = projected.len();
let mut carried = Vec::new();
for (position, (key, &binding)) in keys.iter().zip(&below.keys).enumerate() {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
projected.push(plan.add_expr_at(Expr::Column(binding), ty, span));
projected_names.push(plan.intern(&format!("__domain_{position}")));
carried.push(ColumnBinding::new(
at_index,
u32::try_from(width + position).expect("projection width"),
));
}
let exprs = plan.add_expr_list(&projected);
let names = plan.add_name_list(&projected_names);
let node =
plan.add_node(Node::Project { input: below.node, index: at_index, exprs, names });
Some(Pushed { node, keys: carried, moved: HashMap::new() })
}
Node::Aggregate { input, index: at_index, groups, aggregates } => {
let below = push(plan, input, domain, index, keys, outer)?;
aggregate(plan, below, at_index, groups, aggregates, domain, index, keys)
}
Node::Distinct { input, on } => {
let below = push(plan, input, domain, index, keys, outer)?;
let map = mapping(keys, &below);
let on = if on.is_empty() {
on
} else {
let held = plan.expr_list(on).to_vec();
let mut kept: Vec<ExprRef> =
held.into_iter().map(|expr| remap(plan, expr, &map)).collect();
for (key, &binding) in keys.iter().zip(&below.keys) {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
kept.push(plan.add_expr_at(Expr::Column(binding), ty, span));
}
plan.add_expr_list(&kept)
};
let node = plan.add_node(Node::Distinct { input: below.node, on });
Some(Pushed { node, keys: below.keys, moved: below.moved })
}
Node::Sort { input, keys: order } => {
let below = push(plan, input, domain, index, keys, outer)?;
let map = mapping(keys, &below);
let mut ordering: Vec<SortKey> = Vec::new();
for (key, &binding) in keys.iter().zip(&below.keys) {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
let expr = plan.add_expr_at(Expr::Column(binding), ty, span);
ordering.push(SortKey { expr, descending: false, nulls_first: false });
}
for key in plan.sort_key_list(order).to_vec() {
let expr = remap(plan, key.expr, &map);
ordering.push(SortKey { expr, ..key });
}
let order = plan.add_sort_keys(&ordering);
let node = plan.add_node(Node::Sort { input: below.node, keys: order });
Some(Pushed { node, keys: below.keys, moved: below.moved })
}
Node::Window { input, index: at_index, partition, order, frame, expressions } => {
let below = push(plan, input, domain, index, keys, outer)?;
window(plan, below, at_index, partition, order, frame, expressions, keys)
}
Node::Limit { input, count, offset } => {
let below = push(plan, input, domain, index, keys, outer)?;
limited(plan, below, None, count, offset, keys)
}
Node::TopN { input, keys: order, count, offset } => {
let below = push(plan, input, domain, index, keys, outer)?;
limited(plan, below, Some(order), Some(count), offset, keys)
}
Node::SetOp { left, right, kind, all, index: at_index } => {
let width = walk::outputs(plan, left)?.len();
if width != walk::outputs(plan, right)?.len() {
return None;
}
let one = branch(plan, left, domain, index, keys, outer)?;
let other = branch(plan, right, domain, index, keys, outer)?;
let node =
plan.add_node(Node::SetOp { left: one, right: other, kind, all, index: at_index });
let carried = (0..keys.len())
.map(|position| {
ColumnBinding::new(
at_index,
u32::try_from(width + position).expect("branch width"),
)
})
.collect();
Some(Pushed { node, keys: carried, moved: HashMap::new() })
}
Node::Values { index: at_index, columns, rows } => {
values(plan, at_index, columns, rows, domain, index, keys)
}
Node::TableFunction { index: at_index, function, args, options, settings, columns } => {
lateral(plan, at_index, function, args, options, settings, columns, domain, index, keys)
}
Node::CrossProduct { left, right } => {
let empty = plan.add_expr_list(&[]);
sides(plan, left, right, JoinKind::Inner, empty, domain, index, keys, outer)
}
Node::Join {
left,
right,
kind: kind @ (JoinKind::Inner | JoinKind::Left | JoinKind::Right | JoinKind::Full),
conditions,
..
} => sides(plan, left, right, kind, conditions, domain, index, keys, outer),
_ => None,
}
}
#[allow(clippy::too_many_arguments)]
fn window(
plan: &mut Plan,
below: Pushed,
at_index: u32,
partition: Slice,
order: Slice,
frame: WindowFrame,
expressions: Slice,
keys: &[Key],
) -> Option<Pushed> {
let map = mapping(keys, &below);
let mut divided = Vec::new();
for (key, &binding) in keys.iter().zip(&below.keys) {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
divided.push(plan.add_expr_at(Expr::Column(binding), ty, span));
}
for expr in plan.expr_list(partition).to_vec() {
divided.push(remap(plan, expr, &map));
}
let partition = plan.add_expr_list(÷d);
let mut ordering = Vec::new();
for key in plan.sort_key_list(order).to_vec() {
let expr = remap(plan, key.expr, &map);
ordering.push(SortKey { expr, ..key });
}
let order = plan.add_sort_keys(&ordering);
let frame = WindowFrame {
start: bound(plan, frame.start, &map),
end: bound(plan, frame.end, &map),
..frame
};
let held = plan.expr_list(expressions).to_vec();
let rewritten: Vec<ExprRef> = held.into_iter().map(|expr| remap(plan, expr, &map)).collect();
let expressions = plan.add_expr_list(&rewritten);
let node = plan.add_node(Node::Window {
input: below.node,
index: at_index,
partition,
order,
frame,
expressions,
});
Some(Pushed { node, keys: below.keys, moved: below.moved })
}
fn branch(
plan: &mut Plan,
at: NodeRef,
domain: NodeRef,
index: u32,
keys: &[Key],
outer: &TableSet,
) -> Option<NodeRef> {
let before = walk::outputs(plan, at)?;
let below = push(plan, at, domain, index, keys, outer)?;
let span = plan.expr_span(keys.first()?.expr);
let mut projected = Vec::new();
let mut named = Vec::new();
for (position, (binding, ty)) in before.into_iter().enumerate() {
let moved = below.moved.get(&binding).copied().unwrap_or(binding);
projected.push(plan.add_expr_at(Expr::Column(moved), ty, span));
named.push(plan.intern(&format!("__branch_{position}")));
}
for (position, (key, &binding)) in keys.iter().zip(&below.keys).enumerate() {
let ty = plan.expr_type(key.expr).clone();
projected.push(plan.add_expr_at(Expr::Column(binding), ty, span));
named.push(plan.intern(&format!("__domain_{position}")));
}
let exprs = plan.add_expr_list(&projected);
let names = plan.add_name_list(&named);
let at_index = walk::fresh_index(plan);
Some(plan.add_node(Node::Project { input: below.node, index: at_index, exprs, names }))
}
fn limited(
plan: &mut Plan,
below: Pushed,
order: Option<Slice>,
count: Option<u64>,
offset: u64,
keys: &[Key],
) -> Option<Pushed> {
if count.is_none() && offset == 0 {
return Some(below);
}
let map = mapping(keys, &below);
let span = plan.expr_span(keys.first()?.expr);
let mut divided = Vec::new();
for (key, &binding) in keys.iter().zip(&below.keys) {
let ty = plan.expr_type(key.expr).clone();
divided.push(plan.add_expr_at(Expr::Column(binding), ty, span));
}
let partition = plan.add_expr_list(÷d);
let mut ordering = Vec::new();
for key in order.map(|order| plan.sort_key_list(order).to_vec()).unwrap_or_default() {
let expr = remap(plan, key.expr, &map);
ordering.push(SortKey { expr, ..key });
}
let order = plan.add_sort_keys(&ordering);
let frame = WindowFrame {
unit: WindowUnit::Range,
start: WindowBound::UnboundedPreceding,
end: WindowBound::CurrentRow,
exclude: WindowExclude::NoOthers,
};
let name = plan.intern("row_number");
let args = plan.add_expr_list(&[]);
let call = plan.add_expr_at(
Expr::Window { name, args, distinct: false, filter: None, ignore_nulls: false },
LogicalType::BigInt,
span,
);
let expressions = plan.add_expr_list(&[call]);
let at_index = walk::fresh_index(plan);
let node = plan.add_node(Node::Window {
input: below.node,
index: at_index,
partition,
order,
frame,
expressions,
});
let numbered =
plan.add_expr_at(Expr::Column(ColumnBinding::new(at_index, 0)), LogicalType::BigInt, span);
let mut bounds = Vec::new();
if offset > 0 {
let at = i64::try_from(offset).ok()?;
let value = plan.add_value(Value::BigInt(at));
let right = plan.add_expr_at(Expr::Constant(value), LogicalType::BigInt, span);
bounds.push(plan.add_expr_at(
Expr::Compare { op: CompareOp::Greater, left: numbered, right },
LogicalType::Boolean,
span,
));
}
if let Some(count) = count {
let at = i64::try_from(offset.saturating_add(count)).ok()?;
let value = plan.add_value(Value::BigInt(at));
let right = plan.add_expr_at(Expr::Constant(value), LogicalType::BigInt, span);
bounds.push(plan.add_expr_at(
Expr::Compare { op: CompareOp::LessOrEqual, left: numbered, right },
LogicalType::Boolean,
span,
));
}
let predicate = match bounds.len() {
0 => return None,
1 => bounds[0],
_ => {
let children = plan.add_expr_list(&bounds);
plan.add_expr_at(
Expr::Conjunction { op: ConjunctionOp::And, children },
LogicalType::Boolean,
span,
)
}
};
let node = plan.add_node(Node::Filter { input: node, predicate });
Some(Pushed { node, keys: below.keys, moved: below.moved })
}
fn bound(
plan: &mut Plan,
at: WindowBound,
map: &HashMap<ColumnBinding, ColumnBinding>,
) -> WindowBound {
match at {
WindowBound::Preceding(expr) => WindowBound::Preceding(remap(plan, expr, map)),
WindowBound::Following(expr) => WindowBound::Following(remap(plan, expr, map)),
other => other,
}
}
fn values(
plan: &mut Plan,
at_index: u32,
columns: Slice,
rows: Slice,
domain: NodeRef,
index: u32,
keys: &[Key],
) -> Option<Pushed> {
let held: Vec<Vec<ExprRef>> =
plan.row_list(rows).to_vec().into_iter().map(|row| plan.expr_list(row).to_vec()).collect();
let fields = plan.field_list(columns).to_vec();
if held.is_empty() {
return None;
}
let span = plan.expr_span(keys.first()?.expr);
let mut map = HashMap::new();
for (position, key) in keys.iter().enumerate() {
let at = u32::try_from(position).expect("key count");
map.insert(key.binding, ColumnBinding::new(index, at));
}
let (input, chosen) = if held.len() == 1 {
let row = held.into_iter().next().expect("one row");
(domain, row.into_iter().map(|expr| remap(plan, expr, &map)).collect::<Vec<_>>())
} else {
let counter = LogicalType::Integer;
let count = i32::try_from(held.len()).ok()?;
let numbers: Vec<ExprRef> = (0..count)
.map(|at| {
let value = plan.add_value(Value::Integer(at));
plan.add_expr_at(Expr::Constant(value), counter.clone(), span)
})
.collect();
let numbered: Vec<Slice> =
numbers.iter().map(|&expr| plan.add_expr_list(&[expr])).collect();
let counted = plan.add_rows(&numbered);
let field = Field { name: "__row".to_string(), ty: counter.clone(), not_null: true };
let named = plan.add_fields(&[field]);
let table = walk::fresh_index(plan);
let source = plan.add_node(Node::Values { index: table, columns: named, rows: counted });
let input = plan.add_node(Node::CrossProduct { left: domain, right: source });
let which = plan.add_expr_at(Expr::Column(ColumnBinding::new(table, 0)), counter, span);
let mut chosen = Vec::new();
for (column, field) in fields.iter().enumerate() {
let (last, rest) = held.split_last().expect("more than one row");
let mut arms = Vec::new();
for (at, written) in rest.iter().enumerate() {
let when = plan.add_expr_at(
Expr::Compare { op: CompareOp::Equal, left: which, right: numbers[at] },
LogicalType::Boolean,
span,
);
let then = remap(plan, written[column], &map);
arms.push(Arm { when, then });
}
let otherwise = remap(plan, last[column], &map);
let arms = plan.add_arms(&arms);
chosen.push(plan.add_expr_at(
Expr::Case { arms, otherwise: Some(otherwise) },
field.ty.clone(),
span,
));
}
(input, chosen)
};
let mut projected = chosen;
let mut names: Vec<_> = fields.iter().map(|field| plan.intern(&field.name)).collect();
let width = projected.len();
let mut carried = Vec::new();
for (position, key) in keys.iter().enumerate() {
let ty = plan.expr_type(key.expr).clone();
let at = u32::try_from(position).expect("key count");
projected.push(plan.add_expr_at(Expr::Column(ColumnBinding::new(index, at)), ty, span));
names.push(plan.intern(&format!("__domain_{position}")));
carried.push(ColumnBinding::new(at_index, u32::try_from(width + position).expect("width")));
}
let exprs = plan.add_expr_list(&projected);
let names = plan.add_name_list(&names);
let node = plan.add_node(Node::Project { input, index: at_index, exprs, names });
Some(Pushed { node, keys: carried, moved: HashMap::new() })
}
#[allow(clippy::too_many_arguments)]
fn lateral(
plan: &mut Plan,
at_index: u32,
function: StrRef,
args: Slice,
options: Slice,
settings: Slice,
columns: Slice,
domain: NodeRef,
index: u32,
keys: &[Key],
) -> Option<Pushed> {
let mut map = HashMap::new();
for (position, key) in keys.iter().enumerate() {
let at = u32::try_from(position).expect("key count");
map.insert(key.binding, ColumnBinding::new(index, at));
}
let held = plan.expr_list(args).to_vec();
let rewritten: Vec<ExprRef> = held.into_iter().map(|expr| remap(plan, expr, &map)).collect();
let args = plan.add_expr_list(&rewritten);
let node = plan.add_node(Node::LateralFunction {
input: domain,
index: at_index,
function,
args,
options,
settings,
columns,
});
let carried = (0..keys.len())
.map(|position| ColumnBinding::new(index, u32::try_from(position).expect("key count")))
.collect();
Some(Pushed { node, keys: carried, moved: HashMap::new() })
}
#[allow(clippy::too_many_arguments)]
fn aggregate(
plan: &mut Plan,
below: Pushed,
at_index: u32,
groups: Slice,
aggregates: Slice,
domain: NodeRef,
index: u32,
keys: &[Key],
) -> Option<Pushed> {
let map = mapping(keys, &below);
let held = plan.expr_list(groups).to_vec();
let was_grouped = !held.is_empty();
let mut grouped: Vec<ExprRef> = held.into_iter().map(|expr| remap(plan, expr, &map)).collect();
let held = plan.expr_list(aggregates).to_vec();
let calls: Vec<ExprRef> = held.into_iter().map(|expr| remap(plan, expr, &map)).collect();
let width = grouped.len();
let mut carried = Vec::new();
for (position, (key, &binding)) in keys.iter().zip(&below.keys).enumerate() {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
grouped.push(plan.add_expr_at(Expr::Column(binding), ty, span));
carried.push(ColumnBinding::new(
at_index,
u32::try_from(width + position).expect("grouping width"),
));
}
if was_grouped {
let mut moved = HashMap::new();
for position in 0..calls.len() {
let was = u32::try_from(width + position).expect("aggregate width");
let now = u32::try_from(width + keys.len() + position).expect("aggregate width");
moved.insert(ColumnBinding::new(at_index, was), ColumnBinding::new(at_index, now));
}
let groups = plan.add_expr_list(&grouped);
let aggregates = plan.add_expr_list(&calls);
let node = plan.add_node(Node::Aggregate {
input: below.node,
index: at_index,
groups,
aggregates,
});
return Some(Pushed { node, keys: carried, moved });
}
let inner = walk::fresh_index(plan);
let groups = plan.add_expr_list(&grouped);
let aggregates = plan.add_expr_list(&calls);
let node =
plan.add_node(Node::Aggregate { input: below.node, index: inner, groups, aggregates });
let marker = walk::fresh_index(plan);
let mut carried_exprs = Vec::new();
let mut carried_names = Vec::new();
for (position, key) in keys.iter().enumerate() {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
let at = u32::try_from(position).expect("key count");
carried_exprs.push(plan.add_expr_at(Expr::Column(ColumnBinding::new(inner, at)), ty, span));
carried_names.push(plan.intern(&format!("__domain_{position}")));
}
for (position, &call) in calls.iter().enumerate() {
let ty = plan.expr_type(call).clone();
let span = plan.expr_span(call);
let at = u32::try_from(keys.len() + position).expect("aggregate width");
carried_exprs.push(plan.add_expr_at(Expr::Column(ColumnBinding::new(inner, at)), ty, span));
carried_names.push(plan.intern(&format!("__aggregate_{position}")));
}
let present = plan.add_value(Value::Boolean(true));
carried_exprs.push(plan.add_expr(Expr::Constant(present), LogicalType::Boolean));
carried_names.push(plan.intern("__present"));
let exprs = plan.add_expr_list(&carried_exprs);
let names = plan.add_name_list(&carried_names);
let answered = plan.add_node(Node::Project { input: node, index: marker, exprs, names });
let mut conditions = Vec::new();
for (position, key) in keys.iter().enumerate() {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
let at = u32::try_from(position).expect("key count");
let here = plan.add_expr_at(Expr::Column(ColumnBinding::new(index, at)), ty.clone(), span);
let there = plan.add_expr_at(Expr::Column(ColumnBinding::new(marker, at)), ty, span);
conditions.push(plan.add_expr_at(
Expr::Compare { op: CompareOp::NotDistinctFrom, left: here, right: there },
LogicalType::Boolean,
span,
));
}
let conditions = plan.add_expr_list(&conditions);
let filled = plan.add_node(Node::Join {
left: domain,
right: answered,
kind: JoinKind::Left,
conditions,
build: BuildSide::default(),
});
let present = plan.add_expr(
Expr::Column(ColumnBinding::new(
marker,
u32::try_from(keys.len() + calls.len()).expect("projection width"),
)),
LogicalType::Boolean,
);
let mut repaired = Vec::new();
let mut repaired_names = Vec::new();
for (position, &call) in calls.iter().enumerate() {
let ty = plan.expr_type(call).clone();
let span = plan.expr_span(call);
let at = u32::try_from(keys.len() + position).expect("aggregate width");
let column =
plan.add_expr_at(Expr::Column(ColumnBinding::new(marker, at)), ty.clone(), span);
repaired.push(match empty_answer(plan, call) {
None => column,
Some(value) => {
if ty != value.logical_type() {
return None;
}
let reference = plan.add_value(value);
let otherwise = plan.add_expr_at(Expr::Constant(reference), ty.clone(), span);
let arms = plan.add_arms(&[Arm { when: present, then: column }]);
plan.add_expr_at(Expr::Case { arms, otherwise: Some(otherwise) }, ty, span)
}
});
repaired_names.push(plan.intern(&format!("__aggregate_{position}")));
}
let mut repaired_keys = Vec::new();
for (position, key) in keys.iter().enumerate() {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
let at = u32::try_from(position).expect("key count");
repaired.push(plan.add_expr_at(Expr::Column(ColumnBinding::new(index, at)), ty, span));
repaired_names.push(plan.intern(&format!("__domain_{position}")));
repaired_keys.push(ColumnBinding::new(
at_index,
u32::try_from(calls.len() + position).expect("projection width"),
));
}
let exprs = plan.add_expr_list(&repaired);
let names = plan.add_name_list(&repaired_names);
let node = plan.add_node(Node::Project { input: filled, index: at_index, exprs, names });
Some(Pushed { node, keys: repaired_keys, moved: HashMap::new() })
}
fn empty_answer(plan: &Plan, call: ExprRef) -> Option<Value> {
let Expr::Aggregate { name, .. } = *plan.expr(call) else {
return None;
};
matches!(plan.string(name), "count" | "count_star").then_some(Value::BigInt(0))
}
#[allow(clippy::too_many_arguments)]
fn sides(
plan: &mut Plan,
left: NodeRef,
right: NodeRef,
kind: JoinKind,
conditions: Slice,
domain: NodeRef,
index: u32,
keys: &[Key],
outer: &TableSet,
) -> Option<Pushed> {
let left_reads = correlated(plan, left, outer);
let right_reads = correlated(plan, right, outer);
let keeps_left = matches!(kind, JoinKind::Left | JoinKind::Full);
let keeps_right = matches!(kind, JoinKind::Right | JoinKind::Full);
let carry_right = right_reads || keeps_right;
let carry_left = left_reads || keeps_left || !carry_right;
let (other_domain, other_index) =
if keeps_left && keeps_right { twin(plan, domain)? } else { (domain, index) };
let first = if carry_left { Some(push(plan, left, domain, index, keys, outer)?) } else { None };
let second = if carry_right {
Some(push(plan, right, other_domain, other_index, keys, outer)?)
} else {
None
};
let mut moved = HashMap::new();
for side in [first.as_ref(), second.as_ref()].into_iter().flatten() {
moved.extend(side.moved.clone());
}
let carrier = first.as_ref().or(second.as_ref())?;
let mut map = moved.clone();
for (key, &binding) in keys.iter().zip(&carrier.keys) {
map.insert(key.binding, binding);
}
let held = plan.expr_list(conditions).to_vec();
let mut all: Vec<ExprRef> = held.into_iter().map(|expr| remap(plan, expr, &map)).collect();
if let (Some(one), Some(other)) = (&first, &second) {
for (key, (&here, &there)) in keys.iter().zip(one.keys.iter().zip(&other.keys)) {
let ty = plan.expr_type(key.expr).clone();
let span = plan.expr_span(key.expr);
let mine = plan.add_expr_at(Expr::Column(here), ty.clone(), span);
let yours = plan.add_expr_at(Expr::Column(there), ty, span);
all.push(plan.add_expr_at(
Expr::Compare { op: CompareOp::NotDistinctFrom, left: mine, right: yours },
LogicalType::Boolean,
span,
));
}
}
let conditions = plan.add_expr_list(&all);
let node = plan.add_node(Node::Join {
left: first.as_ref().map_or(left, |one| one.node),
right: second.as_ref().map_or(right, |other| other.node),
kind,
conditions,
build: BuildSide::default(),
});
if keeps_left && keeps_right {
let one = first.as_ref()?.keys.clone();
let other = second.as_ref()?.keys.clone();
return either(plan, node, &one, &other, keys, moved);
}
let carried = if keeps_right { second.as_ref()?.keys.clone() } else { carrier.keys.clone() };
Some(Pushed { node, keys: carried, moved })
}
fn twin(plan: &mut Plan, domain: NodeRef) -> Option<(NodeRef, u32)> {
let Node::Aggregate { input, groups, aggregates, .. } = *plan.node(domain) else {
return None;
};
let index = walk::fresh_index(plan);
Some((plan.add_node(Node::Aggregate { input, index, groups, aggregates }), index))
}
fn either(
plan: &mut Plan,
node: NodeRef,
left: &[ColumnBinding],
right: &[ColumnBinding],
keys: &[Key],
moved: HashMap<ColumnBinding, ColumnBinding>,
) -> Option<Pushed> {
let span = plan.expr_span(keys.first()?.expr);
let produced = walk::outputs(plan, node)?;
let at_index = walk::fresh_index(plan);
let mut exprs = Vec::new();
let mut names = Vec::new();
let mut relabel = HashMap::new();
for (binding, ty) in produced {
if left.contains(&binding) || right.contains(&binding) {
continue;
}
let position = exprs.len();
relabel.insert(binding, at(at_index, position));
exprs.push(plan.add_expr_at(Expr::Column(binding), ty, span));
names.push(plan.intern(&format!("__kept_{position}")));
}
let mut carried = Vec::new();
for (position, (key, (&here, &there))) in keys.iter().zip(left.iter().zip(right)).enumerate() {
let ty = plan.expr_type(key.expr).clone();
let mine = plan.add_expr_at(Expr::Column(here), ty.clone(), span);
let tested = plan.add_expr_at(Expr::Column(here), ty.clone(), span);
let yours = plan.add_expr_at(Expr::Column(there), ty.clone(), span);
let nothing = plan.add_value(Value::Null);
let absent = plan.add_expr_at(Expr::Constant(nothing), ty.clone(), span);
let filled = plan.add_expr_at(
Expr::Compare { op: CompareOp::DistinctFrom, left: tested, right: absent },
LogicalType::Boolean,
span,
);
let arms = plan.add_arms(&[Arm { when: filled, then: mine }]);
carried.push(at(at_index, exprs.len()));
exprs.push(plan.add_expr_at(Expr::Case { arms, otherwise: Some(yours) }, ty, span));
names.push(plan.intern(&format!("__domain_{position}")));
}
let mut told = relabel.clone();
for (before, between) in moved {
if let Some(&now) = relabel.get(&between) {
told.insert(before, now);
}
}
let exprs = plan.add_expr_list(&exprs);
let names = plan.add_name_list(&names);
let node = plan.add_node(Node::Project { input: node, index: at_index, exprs, names });
Some(Pushed { node, keys: carried, moved: told })
}
fn at(index: u32, position: usize) -> ColumnBinding {
ColumnBinding::new(index, u32::try_from(position).expect("a column count fits in a u32"))
}
fn mapping(keys: &[Key], below: &Pushed) -> HashMap<ColumnBinding, ColumnBinding> {
let mut map = below.moved.clone();
for (key, &binding) in keys.iter().zip(&below.keys) {
map.insert(key.binding, binding);
}
map
}
fn remap(plan: &mut Plan, expr: ExprRef, map: &HashMap<ColumnBinding, ColumnBinding>) -> ExprRef {
if map.is_empty() {
return expr;
}
if let Expr::Column(binding) = *plan.expr(expr) {
let Some(&moved) = map.get(&binding) else {
return expr;
};
let ty = plan.expr_type(expr).clone();
let span = plan.expr_span(expr);
return plan.add_expr_at(Expr::Column(moved), ty, span);
}
walk::rebuild(plan, expr, &mut |plan, child| remap(plan, child, map))
}
#[cfg(test)]
mod tests {
use crate::unnest;
use rudb_plan::Plan;
fn correlated(above: &str) -> Plan {
let last = above.lines().next_back().expect("at least one operator above the filter");
let depth = last.len() - last.trim_start().len() + 2;
let inner = " ".repeat(depth);
let deeper = " ".repeat(depth + 2);
let text = format!(
"DependentJoin SINGLE on=[]\n Get memory.main.outer AS o #0 [k::INTEGER]\n{above}{inner}Filter (#1.0::INTEGER = #0.0::INTEGER)::BOOLEAN\n{deeper}Get memory.main.inner AS i #1 [k::INTEGER, value::INTEGER]\n"
);
Plan::parse(&text).expect("a correlated plan")
}
#[test]
fn a_distinct_inside_the_subquery_is_pushed_through() {
let mut plan = correlated(" Distinct on=[]\n Project #2 [#1.1::INTEGER AS value]\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("__domain_0"), "{after}");
assert!(after.contains("CrossProduct"), "{after}");
}
#[test]
fn a_distinct_on_gains_the_domain_columns() {
let mut plan =
correlated(" Distinct on=[#2.0::INTEGER]\n Project #2 [#1.1::INTEGER AS value]\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("Distinct on=[#2.0::INTEGER, #2.1::INTEGER]"), "{after}");
}
#[test]
fn a_sort_inside_the_subquery_orders_by_the_domain_first() {
let mut plan = correlated(
" Sort [#2.0::INTEGER ASC NULLS LAST]\n Project #2 [#1.1::INTEGER AS value]\n",
);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(
after.contains("Sort [#2.1::INTEGER ASC NULLS LAST, #2.0::INTEGER ASC NULLS LAST]"),
"{after}"
);
}
const FRAME: &str = "frame=RANGE UNBOUNDED PRECEDING TO CURRENT ROW EXCLUDE NO OTHERS";
#[test]
fn a_window_inside_the_subquery_partitions_by_the_domain() {
let mut plan = correlated(&format!(
" Window #2 partition=[] order=[#1.1::INTEGER ASC NULLS LAST] {FRAME} expressions=[row_number()::BIGINT]\n"
));
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("partition=[#3.0::INTEGER]"), "{after}");
}
#[test]
fn a_window_that_already_partitions_puts_the_domain_in_front() {
let mut plan = correlated(&format!(
" Window #2 partition=[#1.1::INTEGER] order=[] {FRAME} expressions=[rank()::BIGINT]\n"
));
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("partition=[#3.0::INTEGER, #1.1::INTEGER]"), "{after}");
}
#[test]
fn a_top_n_inside_the_subquery_becomes_a_row_number_per_domain_value() {
let mut plan = correlated(" TopN 1 offset 0 [#1.1::INTEGER ASC NULLS LAST]\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(!after.contains("TopN"), "{after}");
assert!(after.contains("row_number()"), "{after}");
assert!(after.contains("partition=[#2.0::INTEGER]"), "{after}");
assert!(after.contains("<= 1::BIGINT"), "{after}");
}
#[test]
fn an_offset_is_the_lower_bound_and_the_count_is_still_measured_from_the_start() {
let mut plan = correlated(" Limit 2 offset 3\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("> 3::BIGINT"), "{after}");
assert!(after.contains("<= 5::BIGINT"), "{after}");
}
#[test]
fn a_plain_limit_numbers_the_rows_in_whatever_order_they_arrive() {
let mut plan = correlated(" Limit 1 offset 0\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("order=[] "), "{after}");
assert!(after.contains("<= 1::BIGINT"), "{after}");
}
#[test]
fn a_grouped_aggregate_adds_the_domain_to_the_groups() {
let mut plan =
correlated(" Aggregate #2 groups=[#1.1::INTEGER] aggregates=[count_star()::BIGINT]\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("groups=[#1.1::INTEGER, #3.0::INTEGER]"), "{after}");
assert!(!after.contains("__present"), "{after}");
}
#[test]
fn an_ungrouped_count_keeps_its_answer_over_no_rows() {
let mut plan = correlated(" Aggregate #2 groups=[] aggregates=[count_star()::BIGINT]\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("Join LEFT"), "{after}");
assert!(after.contains("TRUE::BOOLEAN AS __present"), "{after}");
assert!(after.contains("CASE WHEN"), "{after}");
assert!(after.contains("ELSE 0::BIGINT"), "{after}");
}
#[test]
fn an_ungrouped_sum_is_left_alone_because_null_is_already_its_answer() {
let mut plan =
correlated(" Aggregate #2 groups=[] aggregates=[sum(#1.1::INTEGER)::HUGEINT]\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("Join LEFT"), "{after}");
assert!(!after.contains("CASE WHEN"), "{after}");
}
#[test]
fn a_limit_over_a_projection_reads_the_domain_the_projection_carried_up() {
let mut plan = correlated(" Limit 1 offset 0\n Project #2 [#1.1::INTEGER AS value]\n");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("partition=[#2.1::INTEGER]"), "{after}");
}
fn correlated_values(rows: &str) -> Plan {
let text = format!(
"DependentJoin INNER on=[]\n Get memory.main.outer AS o #0 [k::INTEGER]\n Values #1 [col0::INTEGER] rows=[{rows}]\n"
);
Plan::parse(&text).expect("a correlated plan")
}
#[test]
fn one_values_row_reading_the_outer_row_becomes_a_projection_over_the_domain() {
let mut plan = correlated_values("[\"*\"(#0.0::INTEGER, 3::INTEGER)::INTEGER]");
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(!after.contains("CrossProduct"), "{after}");
assert!(!after.contains("CASE WHEN"), "{after}");
assert!(after.contains("__domain_0"), "{after}");
}
#[test]
fn several_values_rows_are_chosen_between_by_a_row_number() {
let mut plan = correlated_values(
"[\"*\"(#0.0::INTEGER, 3::INTEGER)::INTEGER], [\"+\"(#0.0::INTEGER, 1::INTEGER)::INTEGER]",
);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(
after.contains("Values #3 [__row::INTEGER] rows=[[0::INTEGER], [1::INTEGER]]"),
"{after}"
);
assert!(after.contains("CASE WHEN (#3.0::INTEGER = 0::INTEGER)"), "{after}");
assert!(after.contains("ELSE \"+\"(#2.0::INTEGER, 1::INTEGER)::INTEGER"), "{after}");
}
fn correlated_set(kind: &str, right: &str) -> Plan {
let text = format!(
"DependentJoin SINGLE on=[]\n Get memory.main.outer AS o #0 [k::INTEGER]\n SetOp {kind} #4\n Project #2 [#1.1::INTEGER AS value]\n Filter (#1.0::INTEGER = #0.0::INTEGER)::BOOLEAN\n Get memory.main.inner AS i #1 [k::INTEGER, value::INTEGER]\n{right}"
);
Plan::parse(&text).expect("a correlated set operation")
}
const PLAIN: &str = " Project #3 [#5.1::INTEGER AS value]\n Get memory.main.other AS u #5 [k::INTEGER, value::INTEGER]\n";
const ALSO: &str = " Project #3 [#5.1::INTEGER AS value]\n Filter (#5.0::INTEGER = #0.0::INTEGER)::BOOLEAN\n Get memory.main.other AS u #5 [k::INTEGER, value::INTEGER]\n";
#[test]
fn a_union_inside_the_subquery_carries_the_domain_out_of_both_sides() {
let mut plan = correlated_set("UNION ALL", PLAIN);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("SetOp UNION ALL"), "{after}");
assert_eq!(after.matches("AS __branch_0, ").count(), 2, "{after}");
assert!(after.contains("CrossProduct"), "{after}");
}
#[test]
fn an_except_inside_the_subquery_subtracts_inside_one_outer_row() {
let mut plan = correlated_set("EXCEPT DISTINCT", ALSO);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("SetOp EXCEPT DISTINCT"), "{after}");
assert_eq!(after.matches("AS __branch_0, ").count(), 2, "{after}");
}
#[test]
fn the_columns_a_side_had_stay_in_front_of_the_domain_columns() {
let mut plan = correlated_set("INTERSECT DISTINCT", PLAIN);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert_eq!(after.matches("AS __branch_0, #").count(), 2, "{after}");
assert!(after.contains("SetOp INTERSECT DISTINCT"), "{after}");
}
fn correlated_join(kind: &str, left: &str, right: &str) -> Plan {
let text = format!(
"DependentJoin SINGLE on=[]\n Get memory.main.outer AS o #0 [k::INTEGER]\n Join {kind} on=[(#1.0::INTEGER = #5.0::INTEGER)::BOOLEAN]\n{left}{right}"
);
Plan::parse(&text).expect("a correlated join")
}
const READS: &str = " Filter (#1.0::INTEGER = #0.0::INTEGER)::BOOLEAN\n Get memory.main.inner AS i #1 [k::INTEGER, value::INTEGER]\n";
const QUIET: &str = " Get memory.main.other AS u #5 [k::INTEGER, value::INTEGER]\n";
#[test]
fn a_right_join_puts_the_domain_on_the_side_it_preserves() {
let mut plan = correlated_join("RIGHT", READS, QUIET);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("Join RIGHT"), "{after}");
assert_eq!(after.matches("CrossProduct").count(), 2, "{after}");
assert_eq!(after.matches("IS NOT DISTINCT FROM").count(), 2, "{after}");
}
#[test]
fn a_right_join_whose_correlated_side_is_the_one_it_preserves_needs_only_that_side() {
let mut plan = correlated_join("RIGHT", QUIET, READS);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("Join RIGHT"), "{after}");
assert_eq!(after.matches("CrossProduct").count(), 1, "{after}");
assert_eq!(after.matches("IS NOT DISTINCT FROM").count(), 1, "{after}");
}
#[test]
fn a_left_join_whose_correlated_side_is_the_right_one_carries_the_domain_on_both() {
let mut plan = correlated_join("LEFT", QUIET, READS);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("Join LEFT"), "{after}");
assert_eq!(after.matches("CrossProduct").count(), 2, "{after}");
assert_eq!(after.matches("IS NOT DISTINCT FROM").count(), 2, "{after}");
}
#[test]
fn a_full_join_reads_whichever_copy_of_the_domain_is_filled_in() {
let mut plan = correlated_join("FULL", READS, QUIET);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("Join FULL"), "{after}");
assert_eq!(after.matches("CrossProduct").count(), 2, "{after}");
assert!(after.contains("CASE WHEN"), "{after}");
assert!(after.contains("IS DISTINCT FROM NULL"), "{after}");
assert!(after.contains("AS __domain_0"), "{after}");
assert_eq!(after.matches("groups=[#0.0::INTEGER] aggregates=[]").count(), 2, "{after}");
let case = after.split("CASE WHEN ").nth(1).expect("a CASE in the plan");
let then = case.split("THEN ").nth(1).expect("a THEN");
let otherwise = case.split("ELSE ").nth(1).expect("an ELSE");
let column = |text: &str| text.split_whitespace().next().expect("a column").to_owned();
assert_ne!(column(then), column(otherwise), "{after}");
}
#[test]
fn a_full_join_keeps_the_columns_it_did_not_replace() {
let mut plan = correlated_join("FULL", QUIET, READS);
unnest::lower(&mut plan).expect("unnesting succeeds");
plan.validate().expect("the rewritten plan is valid");
let after = plan.to_string();
assert!(!after.contains("DependentJoin"), "{after}");
assert!(after.contains("AS __kept_3"), "{after}");
assert!(!after.contains("AS __kept_4"), "{after}");
}
}