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 };
}
}
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;
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()
}
#[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);
}
}