use rudb_common::PhysicalType;
use rudb_plan::{BuildSide, ColumnBinding, Expr, JoinKind, Node, NodeRef, Plan};
use rudb_common::Result;
use crate::estimate::{self, Facts};
use crate::filter;
use crate::pass::{Context, Pass, top_down};
use crate::tables::{Tables, produced};
#[derive(Debug)]
pub struct BuildSideProbeSide;
impl Pass for BuildSideProbeSide {
fn name(&self) -> &'static str {
"build_side_probe_side"
}
fn run(&self, plan: &mut Plan, context: &Context) -> Result<()> {
choose(plan, context.facts());
Ok(())
}
}
#[must_use]
fn prefers(left: u64, right: u64, lookup: bool) -> Option<BuildSide> {
let (bigger, smaller) = if lookup {
(BuildSide::Right, BuildSide::Left)
} else {
(BuildSide::Left, BuildSide::Right)
};
match left.cmp(&right) {
std::cmp::Ordering::Greater => Some(bigger),
std::cmp::Ordering::Less => Some(smaller),
std::cmp::Ordering::Equal => None,
}
}
fn choose(plan: &mut Plan, stats: &Facts) {
let mut tables = Tables::new();
for node in top_down(plan) {
let Node::Join { left, right, kind, conditions, .. } = *plan.node(node) else {
continue;
};
let sides = (left, right);
let turned = matches!(kind, JoinKind::Semi | JoinKind::Anti);
if kind.mirrored().is_none() && !turned {
continue;
}
let below = (produced(plan, left), produced(plan, right));
let held: Vec<_> = plan.expr_list(conditions).to_vec();
let lookup = filter::lookup(plan, &mut tables, &held, &below);
if turned && !lookup {
continue;
}
let (Some(left), Some(right)) =
(estimate::rows(plan, left, stats), estimate::rows(plan, right, stats))
else {
continue;
};
let (left, right) = match (lookup, width(plan, sides.0), width(plan, sides.1)) {
(true, Some(one), Some(other)) => {
(left.saturating_mul(one + HASHED), right.saturating_mul(other + HASHED))
}
_ => (left, right),
};
let Some(wanted) = prefers(left, right, lookup) else {
continue;
};
if let Node::Join { build, .. } = plan.node_mut(node) {
*build = wanted;
}
}
}
const HASHED: u64 = 8;
const INLINE: u64 = 12;
fn width(plan: &Plan, node: NodeRef) -> Option<u64> {
let fields = |columns| -> u64 {
plan.field_list(columns).iter().map(|field| field.ty.physical().size() as u64).sum()
};
let exprs = |list| -> u64 {
plan.expr_list(list)
.iter()
.map(|&expr| {
let header = plan.expr_type(expr).physical().size() as u64;
match *plan.expr(expr) {
Expr::Column(binding) => header + spilled(plan, binding),
_ => header,
}
})
.sum()
};
match *plan.node(node) {
Node::Get { index, columns, .. } => Some(
fields(columns)
+ (0..plan.field_list(columns).len())
.map(|column| {
let column = u32::try_from(column).unwrap_or(u32::MAX);
spilled(plan, ColumnBinding { table: index, column })
})
.sum::<u64>(),
),
Node::Values { columns, .. }
| Node::TableFunction { columns, .. }
| Node::CteScan { columns, .. } => Some(fields(columns)),
Node::Project { exprs: list, .. } => Some(exprs(list)),
Node::Aggregate { groups, aggregates, .. } => Some(exprs(groups) + exprs(aggregates)),
Node::Filter { input, .. }
| Node::Sort { input, .. }
| Node::Limit { input, .. }
| Node::LimitPercent { input, .. }
| Node::TopN { input, .. }
| Node::Distinct { input, .. } => width(plan, input),
Node::Join { left, right, kind, .. }
| Node::LinkJoin { child: left, parent: right, kind, .. } => match kind {
JoinKind::Semi | JoinKind::Anti => width(plan, left),
JoinKind::Mark => Some(width(plan, left)? + 1),
_ => Some(width(plan, left)? + width(plan, right)?),
},
Node::CrossProduct { left, right } => Some(width(plan, left)? + width(plan, right)?),
_ => None,
}
}
fn spilled(plan: &Plan, binding: ColumnBinding) -> u64 {
let Some(field) = (0..u32::try_from(plan.node_count()).unwrap_or(u32::MAX)).find_map(|at| {
match *plan.node(at) {
Node::Get { index, columns, .. } if index == binding.table => {
plan.field_list(columns).get(binding.column as usize)
}
_ => None,
}
}) else {
return 0;
};
if field.ty.physical() != PhysicalType::Varlen {
return 0;
}
plan.width_measured(binding.table, &field.name).filter(|&bytes| bytes > INLINE).unwrap_or(0)
}
#[cfg(test)]
mod tests {
use rudb_plan::{BuildSide, JoinKind, Node, Plan};
use super::{BuildSideProbeSide, prefers};
use crate::estimate::Facts;
use crate::pass::{Context, Pass};
fn joined(kind: &str, left: u64, right: u64) -> (Plan, Context) {
let text = format!(
"Join {kind} on=[]\n Get memory.main.l AS l #0 [a::BIGINT]\n Get memory.main.r AS r #1 [b::BIGINT]\n"
);
let plan =
Plan::parse(&text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
let mut facts = Facts::new();
facts.record("memory", "main", "l", left);
facts.record("memory", "main", "r", right);
let mut context = Context::new();
context.measure(std::sync::Arc::new(facts));
(plan, context)
}
fn chosen(kind: &str, left: u64, right: u64) -> BuildSide {
let (mut plan, context) = joined(kind, left, right);
BuildSideProbeSide.run(&mut plan, &context).expect("the pass does not fail");
let Node::Join { build, .. } = *plan.node(plan.root()) else {
panic!("the root stopped being a join");
};
build
}
fn keyed(kind: &str, left: u64, right: u64) -> BuildSide {
let text = format!(
"Join {kind} on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n Get memory.main.l AS l #0 [a::BIGINT]\n Get memory.main.r AS r #1 [b::BIGINT]\n"
);
side(&text, left, right)
}
fn side(text: &str, left: u64, right: u64) -> BuildSide {
measured(text, left, right, &[])
}
fn measured(text: &str, left: u64, right: u64, widths: &[(u32, &str, u64)]) -> BuildSide {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
for &(index, column, bytes) in widths {
plan.measure_width(index, column, bytes);
}
let mut facts = Facts::new();
facts.record("memory", "main", "l", left);
facts.record("memory", "main", "r", right);
let mut context = Context::new();
context.measure(std::sync::Arc::new(facts));
BuildSideProbeSide.run(&mut plan, &context).expect("the pass does not fail");
let Node::Join { build, .. } = *plan.node(plan.root()) else {
panic!("the root stopped being a join");
};
build
}
#[test]
fn the_larger_side_is_the_one_gathered_because_the_nested_loop_walks_it_a_chunk_at_a_time() {
assert_eq!(prefers(100_000, 4, false), Some(BuildSide::Left));
assert_eq!(prefers(4, 100_000, false), Some(BuildSide::Right));
}
#[test]
fn the_smaller_side_is_the_one_gathered_when_a_hash_table_is_going_to_be_built_out_of_it() {
assert_eq!(prefers(100_000, 4, true), Some(BuildSide::Right));
assert_eq!(prefers(4, 100_000, true), Some(BuildSide::Left));
}
#[test]
fn two_sides_of_the_same_size_are_not_a_preference() {
assert_eq!(prefers(500, 500, false), None);
assert_eq!(prefers(500, 500, true), None);
}
#[test]
fn a_join_with_an_equality_gathers_the_small_side_and_the_same_join_without_one_does_not() {
assert_eq!(keyed("INNER", 400_000, 4), BuildSide::Right);
assert_eq!(chosen("INNER", 400_000, 4), BuildSide::Left);
}
#[test]
fn a_hash_join_gathers_the_side_with_fewer_bytes_even_when_it_has_more_rows() {
let text = "Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n Get memory.main.l AS l #0 [a::BIGINT]\n Get memory.main.r AS r #1 [b::BIGINT, c::VARCHAR, d::VARCHAR, e::VARCHAR, f::VARCHAR, g::VARCHAR]\n";
assert_eq!(side(text, 100_000, 40_000), BuildSide::Left);
assert_eq!(keyed("INNER", 100_000, 40_000), BuildSide::Right);
}
#[test]
fn a_string_longer_than_its_header_holds_counts_its_bytes_as_well() {
let text = "Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n Get memory.main.l AS l #0 [a::BIGINT, n::VARCHAR]\n Get memory.main.r AS r #1 [b::BIGINT, c::VARCHAR]\n";
assert_eq!(measured(text, 100_000, 30_000, &[]), BuildSide::Right);
assert_eq!(measured(text, 100_000, 30_000, &[(1, "c", 100)]), BuildSide::Left);
assert_eq!(measured(text, 100_000, 30_000, &[(1, "c", 12)]), BuildSide::Right);
let projected = "Join INNER on=[(#0.0::BIGINT = #2.0::BIGINT)::BOOLEAN]\n Get memory.main.l AS l #0 [a::BIGINT, n::VARCHAR]\n Project #2 [#1.0::BIGINT AS b, #1.1::VARCHAR AS c]\n Get memory.main.r AS r #1 [b::BIGINT, c::VARCHAR]\n";
assert_eq!(measured(projected, 100_000, 30_000, &[(1, "c", 100)]), BuildSide::Left);
}
#[test]
fn the_nested_loop_still_compares_rows_whatever_the_widths() {
let text = "Join INNER on=[(#0.0::BIGINT < #1.0::BIGINT)::BOOLEAN]\n Get memory.main.l AS l #0 [a::BIGINT]\n Get memory.main.r AS r #1 [b::BIGINT, c::VARCHAR, d::VARCHAR, e::VARCHAR, f::VARCHAR, g::VARCHAR]\n";
assert_eq!(side(text, 100_000, 40_000), BuildSide::Left);
}
#[test]
fn a_condition_no_lookup_answers_still_gathers_the_larger_side() {
let text = "Join INNER on=[(#0.0::BIGINT < #1.0::BIGINT)::BOOLEAN]\n Get memory.main.l AS l #0 [a::BIGINT]\n Get memory.main.r AS r #1 [b::BIGINT]\n";
assert_eq!(side(text, 400_000, 4), BuildSide::Left);
}
#[test]
fn a_big_left_and_a_small_right_gathers_the_left() {
assert_eq!(chosen("INNER", 400_000, 4), BuildSide::Left);
}
#[test]
fn a_small_left_and_a_big_right_keeps_the_side_the_binder_emitted() {
assert_eq!(chosen("INNER", 4, 400_000), BuildSide::Right);
}
#[test]
fn an_outer_join_whose_kind_has_a_mirror_still_gets_the_larger_side() {
assert_eq!(chosen("LEFT", 400_000, 4), BuildSide::Left);
assert_eq!(chosen("FULL", 400_000, 4), BuildSide::Left);
}
#[test]
fn a_left_join_on_a_key_gathers_whichever_side_is_smaller() {
assert_eq!(keyed("LEFT", 4, 400_000), BuildSide::Left);
assert_eq!(keyed("LEFT", 400_000, 4), BuildSide::Right);
}
#[test]
fn a_right_join_on_a_key_is_the_same_rule_the_other_way_round() {
assert_eq!(keyed("RIGHT", 400_000, 4), BuildSide::Right);
assert_eq!(keyed("RIGHT", 4, 400_000), BuildSide::Left);
}
#[test]
fn a_full_join_on_a_key_still_gathers_the_smaller_side() {
assert_eq!(keyed("FULL", 4, 400_000), BuildSide::Left);
assert_eq!(keyed("FULL", 400_000, 4), BuildSide::Right);
}
#[test]
fn a_kind_that_names_its_left_input_as_the_subject_is_left_alone() {
assert_eq!(chosen("SEMI", 400_000, 4), BuildSide::Right);
assert_eq!(chosen("ANTI", 400_000, 4), BuildSide::Right);
assert_eq!(chosen("MARK", 400_000, 4), BuildSide::Right);
assert_eq!(chosen("POSITIONAL", 400_000, 4), BuildSide::Right);
}
#[test]
fn a_table_nobody_measured_leaves_the_join_as_it_was() {
let text = "Join INNER on=[]\n Get memory.main.l AS l #0 [a::BIGINT]\n Get memory.main.r AS r #1 [b::BIGINT]\n";
let mut plan = Plan::parse(text).expect("the plan parses");
let mut facts = Facts::new();
facts.record("memory", "main", "l", 400_000);
let mut context = Context::new();
context.measure(std::sync::Arc::new(facts));
BuildSideProbeSide.run(&mut plan, &context).expect("the pass does not fail");
let Node::Join { build, .. } = *plan.node(plan.root()) else {
panic!("the root stopped being a join");
};
assert_eq!(build, BuildSide::Right, "one estimate is not two estimates");
}
#[test]
fn running_it_twice_writes_what_is_already_there() {
let (mut plan, context) = joined("INNER", 400_000, 4);
BuildSideProbeSide.run(&mut plan, &context).expect("the pass does not fail");
let once = plan.to_string();
BuildSideProbeSide.run(&mut plan, &context).expect("the pass does not fail");
assert_eq!(plan.to_string(), once, "the pass did not settle");
}
#[test]
fn the_kinds_with_no_mirror_are_the_ones_this_pass_refuses_to_touch() {
for kind in
[JoinKind::Semi, JoinKind::Anti, JoinKind::Single, JoinKind::Mark, JoinKind::Positional]
{
assert_eq!(kind.mirrored(), None, "{kind:?} grew a mirror and this test did not");
}
for kind in [JoinKind::Inner, JoinKind::Left, JoinKind::Right, JoinKind::Full] {
assert!(kind.mirrored().is_some(), "{kind:?} lost its mirror");
}
}
}