use std::collections::VecDeque;
use rudb_common::Result;
use rudb_plan::{Expr, JoinKind, Node, NodeRef, Plan};
use crate::pass::{Context, Pass, top_down};
use crate::tables::produced;
use crate::walk;
#[derive(Debug, Clone, Copy)]
pub struct MarkToSemi;
impl Pass for MarkToSemi {
fn name(&self) -> &'static str {
"mark_to_semi"
}
fn run(&self, plan: &mut Plan, _context: &Context) -> Result<()> {
convert(plan);
Ok(())
}
}
pub fn convert(plan: &mut Plan) {
for node in top_down(plan) {
let Node::Filter { input, predicate } = *plan.node(node) else {
continue;
};
let Node::Join { left, right, kind: JoinKind::Mark, conditions, build } = *plan.node(input)
else {
continue;
};
let Expr::Column(tested) = *plan.expr(predicate) else {
continue;
};
let Some(outputs) = walk::outputs(plan, right) else {
continue;
};
let Some((marker, _)) = outputs.last() else {
continue;
};
if tested != *marker || read_above(plan, right, node, input) {
continue;
}
*plan.node_mut(node) = Node::Join { left, right, kind: JoinKind::Semi, conditions, build };
}
}
#[derive(Debug, Clone, Copy)]
pub struct DistinctToSemi;
impl Pass for DistinctToSemi {
fn name(&self) -> &'static str {
"distinct_to_semi"
}
fn run(&self, plan: &mut Plan, _context: &Context) -> Result<()> {
narrow(plan);
Ok(())
}
}
pub fn narrow(plan: &mut Plan) {
for node in top_down(plan) {
let Node::Aggregate { input, groups, aggregates, .. } = *plan.node(node) else {
continue;
};
if !plan.expr_list(aggregates).is_empty() {
continue;
}
let Node::Join { left, right, kind: JoinKind::Inner, conditions, build } =
*plan.node(input)
else {
continue;
};
let driving = produced(plan, left);
let grouped = plan.expr_list(groups).iter().all(|&group| match *plan.expr(group) {
Expr::Column(binding) => driving.contains(binding.table),
_ => false,
});
if !grouped || parents(plan, input) != 1 || read_above(plan, right, node, input) {
continue;
}
*plan.node_mut(input) = Node::Join { left, right, kind: JoinKind::Semi, conditions, build };
}
}
fn parents(plan: &Plan, at: NodeRef) -> usize {
let mut found = 0;
for node in 0..u32::try_from(plan.node_count()).unwrap_or(u32::MAX) {
found +=
plan.node(node).children().into_iter().flatten().filter(|&child| child == at).count();
}
found
}
fn read_above(plan: &Plan, right: NodeRef, filter: NodeRef, join: NodeRef) -> bool {
let gathered = produced(plan, right);
let inside = subtree(plan, right);
let mut found = false;
for at in top_down(plan) {
if at == filter || at == join || inside.contains(&at) {
continue;
}
walk::node_columns(plan, at, &mut |_, binding| {
found |= gathered.contains(binding.table);
});
}
found
}
fn subtree(plan: &Plan, at: NodeRef) -> Vec<NodeRef> {
let mut found = Vec::new();
let mut pending = VecDeque::from([at]);
while let Some(node) = pending.pop_front() {
if found.contains(&node) {
continue;
}
found.push(node);
pending.extend(plan.node(node).children().into_iter().flatten());
}
found
}
#[cfg(test)]
mod tests {
use rudb_plan::Plan;
use super::{convert, narrow};
fn converted(text: &str) -> String {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
convert(&mut plan);
plan.validate().unwrap_or_else(|error| panic!("{text} did not stay valid: {error}"));
plan.to_string()
}
fn narrowed(text: &str) -> String {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
narrow(&mut plan);
plan.validate().unwrap_or_else(|error| panic!("{text} did not stay valid: {error}"));
plan.to_string()
}
#[test]
fn a_filter_on_the_marker_becomes_the_join_itself() {
assert_eq!(
converted(concat!(
"Filter #1.1::BOOLEAN\n",
" Join MARK on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Project #1 [#2.0::BIGINT AS k, TRUE::BOOLEAN AS mark]\n",
" Get memory.main.u AS u #2 [k::BIGINT]\n",
)),
concat!(
"Join SEMI on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Project #1 [#2.0::BIGINT AS k, TRUE::BOOLEAN AS mark]\n",
" Get memory.main.u AS u #2 [k::BIGINT]\n",
)
);
}
#[test]
fn a_filter_on_something_other_than_the_marker_is_left_alone() {
let text = concat!(
"Filter (#0.0::BIGINT > 3::BIGINT)::BOOLEAN\n",
" Join MARK on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Project #1 [#2.0::BIGINT AS k, TRUE::BOOLEAN AS mark]\n",
" Get memory.main.u AS u #2 [k::BIGINT]\n",
);
assert_eq!(converted(text), text);
}
#[test]
fn a_filter_on_a_gathered_column_that_is_not_the_marker_is_left_alone() {
let text = concat!(
"Filter (#1.0::BIGINT > 3::BIGINT)::BOOLEAN\n",
" Join MARK on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Project #1 [#2.0::BIGINT AS k, TRUE::BOOLEAN AS mark]\n",
" Get memory.main.u AS u #2 [k::BIGINT]\n",
);
assert_eq!(converted(text), text);
}
#[test]
fn a_gathered_column_read_above_the_filter_stops_the_rewrite() {
let text = concat!(
"Project #3 [#1.0::BIGINT AS k]\n",
" Filter #1.1::BOOLEAN\n",
" Join MARK on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Project #1 [#2.0::BIGINT AS k, TRUE::BOOLEAN AS mark]\n",
" Get memory.main.u AS u #2 [k::BIGINT]\n",
);
assert_eq!(converted(text), text);
}
#[test]
fn a_driving_column_read_above_the_filter_does_not_stop_it() {
assert_eq!(
converted(concat!(
"Project #3 [#0.0::BIGINT AS a]\n",
" Filter #1.1::BOOLEAN\n",
" Join MARK on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Project #1 [#2.0::BIGINT AS k, TRUE::BOOLEAN AS mark]\n",
" Get memory.main.u AS u #2 [k::BIGINT]\n",
)),
concat!(
"Project #3 [#0.0::BIGINT AS a]\n",
" Join SEMI on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Project #1 [#2.0::BIGINT AS k, TRUE::BOOLEAN AS mark]\n",
" Get memory.main.u AS u #2 [k::BIGINT]\n",
)
);
}
#[test]
fn a_filter_over_a_join_that_is_not_a_mark_is_left_alone() {
let text = concat!(
"Filter #1.1::BOOLEAN\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Project #1 [#2.0::BIGINT AS k, TRUE::BOOLEAN AS mark]\n",
" Get memory.main.u AS u #2 [k::BIGINT]\n",
);
assert_eq!(converted(text), text);
}
#[test]
fn a_duplicate_eliminator_over_an_inner_join_makes_the_join_a_semi_join() {
assert_eq!(
narrowed(concat!(
"Aggregate #3 groups=[#0.0::BIGINT] aggregates=[]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Get memory.main.u AS u #1 [k::BIGINT]\n",
)),
concat!(
"Aggregate #3 groups=[#0.0::BIGINT] aggregates=[]\n",
" Join SEMI on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Get memory.main.u AS u #1 [k::BIGINT]\n",
)
);
}
#[test]
fn an_aggregate_that_actually_aggregates_is_left_alone() {
let text = concat!(
"Aggregate #3 groups=[#0.0::BIGINT] aggregates=[count_star()::BIGINT]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Get memory.main.u AS u #1 [k::BIGINT]\n",
);
assert_eq!(narrowed(text), text);
}
#[test]
fn a_group_that_reads_the_gathered_side_is_left_alone() {
let text = concat!(
"Aggregate #3 groups=[#0.0::BIGINT, #1.0::BIGINT] aggregates=[]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Get memory.main.u AS u #1 [k::BIGINT]\n",
);
assert_eq!(narrowed(text), text);
}
#[test]
fn a_group_over_an_expression_rather_than_a_column_is_left_alone() {
let text = concat!(
"Aggregate #3 groups=[\"+\"(#0.0::BIGINT, 1::BIGINT)::BIGINT] aggregates=[]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Get memory.main.u AS u #1 [k::BIGINT]\n",
);
assert_eq!(narrowed(text), text);
}
#[test]
fn a_join_a_second_node_also_points_at_is_left_alone() {
let text = concat!(
"Join INNER on=[(#0.0::BIGINT = #4.0::BIGINT)::BOOLEAN]\n",
" Aggregate #3 groups=[#0.0::BIGINT] aggregates=[]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Get memory.main.u AS u #1 [k::BIGINT]\n",
" Project #4 [#0.0::BIGINT AS a]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Get memory.main.u AS u #1 [k::BIGINT]\n",
);
assert_eq!(narrowed(text), text);
}
#[test]
fn a_gathered_column_read_somewhere_else_stops_the_narrowing() {
let text = concat!(
"Project #4 [#3.0::BIGINT AS a, #1.0::BIGINT AS k]\n",
" Aggregate #3 groups=[#0.0::BIGINT] aggregates=[]\n",
" Join INNER on=[(#0.0::BIGINT = #1.0::BIGINT)::BOOLEAN]\n",
" Get memory.main.t AS t #0 [a::BIGINT, b::BIGINT]\n",
" Get memory.main.u AS u #1 [k::BIGINT]\n",
);
assert_eq!(narrowed(text), text);
}
}