use std::collections::HashMap;
use rudb_common::{LogicalType, Result, Value};
use rudb_kernels::{Comparison, Connective, call_values, cast_value, combine, compare_values};
use rudb_plan::{CompareOp, ConjunctionOp, Expr, ExprRef, Node, NodeRef, Plan, Slice, SortKey};
use rudb_vector::Vector;
use crate::pass::{Context, Pass, top_down};
use crate::walk;
pub const VOLATILE: [&str; 17] = [
"current_connection_id",
"current_query",
"current_query_id",
"current_transaction_id",
"currval",
"error",
"gen_random_uuid",
"nextval",
"random",
"setseed",
"setval",
"sleep_ms",
"stats",
"uuid",
"uuidv4",
"uuidv7",
"write_log",
];
#[derive(Debug, Clone, Copy)]
pub struct ExpressionRewriter;
impl Pass for ExpressionRewriter {
fn name(&self) -> &'static str {
"expression_rewriter"
}
fn run(&self, plan: &mut Plan, _context: &Context) -> Result<()> {
rewrite(plan);
Ok(())
}
}
type Done = HashMap<ExprRef, ExprRef>;
fn rewrite(plan: &mut Plan) {
let mut done = Done::new();
for node in top_down(plan) {
node_expressions(plan, node, &mut done);
}
}
fn node_expressions(plan: &mut Plan, node: NodeRef, done: &mut Done) {
match *plan.node(node) {
Node::Get { .. }
| Node::Dummy
| Node::Limit { .. }
| Node::SetOp { .. }
| Node::CrossProduct { .. } => {}
Node::Values { rows, .. } => {
let held = plan.row_list(rows).to_vec();
let rewritten: Vec<Slice> =
held.iter().map(|&row| expr_list(plan, row, done).unwrap_or(row)).collect();
if rewritten != held {
let rows = plan.add_rows(&rewritten);
match plan.node_mut(node) {
Node::Values { rows: held, .. } => *held = rows,
_ => unreachable!("the node was a values list a moment ago"),
}
}
}
Node::TableFunction { args, .. } => {
if let Some(rewritten) = expr_list(plan, args, done) {
match plan.node_mut(node) {
Node::TableFunction { args, .. } => *args = rewritten,
_ => unreachable!("the node was a table function a moment ago"),
}
}
}
Node::Filter { predicate, .. } => {
let rewritten = expression(plan, predicate, done);
if rewritten != predicate {
match plan.node_mut(node) {
Node::Filter { predicate, .. } => *predicate = rewritten,
_ => unreachable!("the node was a filter a moment ago"),
}
}
}
Node::Project { exprs, .. } => {
if let Some(rewritten) = expr_list(plan, exprs, done) {
match plan.node_mut(node) {
Node::Project { exprs, .. } => *exprs = rewritten,
_ => unreachable!("the node was a projection a moment ago"),
}
}
}
Node::Aggregate { groups, aggregates, .. } => {
let rewritten_groups = expr_list(plan, groups, done);
let rewritten_aggregates = expr_list(plan, aggregates, done);
match plan.node_mut(node) {
Node::Aggregate { groups, aggregates, .. } => {
if let Some(rewritten) = rewritten_groups {
*groups = rewritten;
}
if let Some(rewritten) = rewritten_aggregates {
*aggregates = rewritten;
}
}
_ => unreachable!("the node was an aggregate a moment ago"),
}
}
Node::Sort { keys, .. } | Node::TopN { keys, .. } => {
let held = plan.sort_key_list(keys).to_vec();
let rewritten: Vec<SortKey> = held
.iter()
.map(|key| SortKey { expr: expression(plan, key.expr, done), ..*key })
.collect();
if rewritten != held {
let keys = plan.add_sort_keys(&rewritten);
match plan.node_mut(node) {
Node::Sort { keys: held, .. } | Node::TopN { keys: held, .. } => *held = keys,
_ => unreachable!("the node was a sort a moment ago"),
}
}
}
Node::Distinct { on, .. } => {
if let Some(rewritten) = expr_list(plan, on, done) {
match plan.node_mut(node) {
Node::Distinct { on, .. } => *on = rewritten,
_ => unreachable!("the node was a distinct a moment ago"),
}
}
}
Node::Join { conditions, .. } => {
if let Some(rewritten) = expr_list(plan, conditions, done) {
match plan.node_mut(node) {
Node::Join { conditions, .. } => *conditions = rewritten,
_ => unreachable!("the node was a join a moment ago"),
}
}
}
}
}
fn expr_list(plan: &mut Plan, slice: Slice, done: &mut Done) -> Option<Slice> {
walk::list(plan, slice, &mut |plan, expr| expression(plan, expr, done))
}
fn expression(plan: &mut Plan, expr: ExprRef, done: &mut Done) -> ExprRef {
if let Some(&already) = done.get(&expr) {
return already;
}
let rebuilt = walk::rebuild(plan, expr, &mut |plan, child| expression(plan, child, done));
let simplified = simplify(plan, rebuilt);
done.insert(expr, simplified);
simplified
}
fn simplify(plan: &mut Plan, expr: ExprRef) -> ExprRef {
if matches!(*plan.expr(expr), Expr::Constant(_)) {
return expr;
}
if let Some(value) = fold(plan, expr) {
if let Some(folded) = constant_of(plan, expr, value) {
return folded;
}
}
match *plan.expr(expr) {
Expr::Conjunction { op, children } => conjunction(plan, expr, op, children),
Expr::Case { arms, otherwise } => case(plan, expr, arms, otherwise),
Expr::Compare { op, left, right } => null_comparison(plan, expr, op, left, right),
_ => expr,
}
}
fn fold(plan: &Plan, expr: ExprRef) -> Option<Value> {
match *plan.expr(expr) {
Expr::Cast { input, try_cast } => {
let inner = constant(plan, input)?;
cast_value(&inner, plan.expr_type(expr), try_cast).ok()
}
Expr::Compare { op, left, right } => {
let left = constant(plan, left)?;
let right = constant(plan, right)?;
compare_values(comparison(op), &left, &right).ok()
}
Expr::Conjunction { op, children } => {
let values = constants(plan, children)?;
let vectors: Vec<Vector> = values
.into_iter()
.map(|value| Vector::constant(LogicalType::Boolean, value, 1))
.collect();
Some(combine(connective(op), &vectors).ok()?.value_at(0))
}
Expr::Function { name, args } => {
let name = plan.string(name);
if VOLATILE.contains(&name) {
return None;
}
let values = constants(plan, args)?;
call_values(name, &values, plan.expr_type(expr)).ok()
}
_ => None,
}
}
fn constant(plan: &Plan, expr: ExprRef) -> Option<Value> {
match *plan.expr(expr) {
Expr::Constant(value) => Some(plan.value(value).clone()),
_ => None,
}
}
fn constants(plan: &Plan, slice: Slice) -> Option<Vec<Value>> {
plan.expr_list(slice).iter().map(|&expr| constant(plan, expr)).collect()
}
fn constant_of(plan: &mut Plan, expr: ExprRef, value: Value) -> Option<ExprRef> {
let ty = plan.expr_type(expr).clone();
if !value.is_null() && value.logical_type() != ty {
return None;
}
let held = plan.add_value(value);
Some(plan.add_expr(Expr::Constant(held), ty))
}
fn conjunction(plan: &mut Plan, expr: ExprRef, op: ConjunctionOp, children: Slice) -> ExprRef {
let (decides, drops) = match op {
ConjunctionOp::And => (false, true),
ConjunctionOp::Or => (true, false),
};
let held = plan.expr_list(children).to_vec();
let mut kept = Vec::with_capacity(held.len());
for child in held.iter().copied() {
match constant(plan, child).as_ref().and_then(Value::as_bool) {
Some(known) if known == decides => {
return constant_of(plan, expr, Value::Boolean(decides)).unwrap_or(expr);
}
Some(_) => {}
None => kept.push(child),
}
}
if kept.len() == held.len() {
return expr;
}
match kept.as_slice() {
[] => constant_of(plan, expr, Value::Boolean(drops)).unwrap_or(expr),
[only] => *only,
rest => {
let children = plan.add_expr_list(rest);
plan.add_expr(Expr::Conjunction { op, children }, LogicalType::Boolean)
}
}
}
enum Fires {
Always,
Never,
Maybe,
}
fn fires(plan: &Plan, when: ExprRef) -> Fires {
match constant(plan, when) {
Some(Value::Boolean(true)) => Fires::Always,
Some(Value::Boolean(false) | Value::Null) => Fires::Never,
Some(_) | None => Fires::Maybe,
}
}
fn case(plan: &mut Plan, expr: ExprRef, arms: Slice, otherwise: Option<ExprRef>) -> ExprRef {
let held = plan.arm_list(arms).to_vec();
let mut kept = Vec::with_capacity(held.len());
let mut result = otherwise;
let mut cut = false;
for arm in held.iter().copied() {
match fires(plan, arm.when) {
Fires::Never => {}
Fires::Always => {
result = Some(arm.then);
cut = true;
break;
}
Fires::Maybe => kept.push(arm),
}
}
if kept.len() == held.len() && !cut {
return expr;
}
if kept.is_empty() {
return match result {
Some(only) if plan.expr_type(only) == plan.expr_type(expr) => only,
Some(_) => expr,
None => constant_of(plan, expr, Value::Null).unwrap_or(expr),
};
}
let ty = plan.expr_type(expr).clone();
let arms = plan.add_arms(&kept);
plan.add_expr(Expr::Case { arms, otherwise: result }, ty)
}
fn null_comparison(
plan: &mut Plan,
expr: ExprRef,
op: CompareOp,
left: ExprRef,
right: ExprRef,
) -> ExprRef {
if matches!(op, CompareOp::DistinctFrom | CompareOp::NotDistinctFrom) {
return expr;
}
let is_null = |side| constant(plan, side).is_some_and(|value| value.is_null());
if is_null(left) || is_null(right) {
constant_of(plan, expr, Value::Null).unwrap_or(expr)
} else {
expr
}
}
fn comparison(op: CompareOp) -> Comparison {
match op {
CompareOp::Equal => Comparison::Equal,
CompareOp::NotEqual => Comparison::NotEqual,
CompareOp::Less => Comparison::Less,
CompareOp::LessOrEqual => Comparison::LessOrEqual,
CompareOp::Greater => Comparison::Greater,
CompareOp::GreaterOrEqual => Comparison::GreaterOrEqual,
CompareOp::DistinctFrom => Comparison::DistinctFrom,
CompareOp::NotDistinctFrom => Comparison::NotDistinctFrom,
}
}
fn connective(op: ConjunctionOp) -> Connective {
match op {
ConjunctionOp::And => Connective::And,
ConjunctionOp::Or => Connective::Or,
}
}
#[cfg(test)]
mod tests {
use super::{ExpressionRewriter, VOLATILE};
use crate::pass::{Context, Pass};
use rudb_plan::Plan;
fn folded(text: &str) -> String {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
ExpressionRewriter
.run(&mut plan, &Context::new())
.unwrap_or_else(|error| panic!("{text} did not fold: {error}"));
plan.validate().unwrap_or_else(|error| panic!("{text} folded to a bad plan: {error}"));
plan.to_string()
}
const SCAN: &str = " Get memory.main.t AS t #0 [a::INTEGER, b::VARCHAR, c::BOOLEAN]\n";
#[test]
fn arithmetic_over_constants_becomes_the_number() {
let before = format!("Project #1 [\"+\"(2::INTEGER, 3::INTEGER)::INTEGER AS n]\n{SCAN}");
let after = format!("Project #1 [5::INTEGER AS n]\n{SCAN}");
assert_eq!(folded(&before), after);
}
#[test]
fn a_nest_of_constants_folds_all_the_way_up_in_one_walk() {
let before = format!(
"Project #1 [\"+\"(\"+\"(1::INTEGER, 2::INTEGER)::INTEGER, 3::INTEGER)::INTEGER AS n]\n{SCAN}"
);
let after = format!("Project #1 [6::INTEGER AS n]\n{SCAN}");
assert_eq!(folded(&before), after);
}
#[test]
fn a_call_with_a_column_in_it_is_left_alone() {
let text = format!("Project #1 [\"+\"(#0.0::INTEGER, 3::INTEGER)::INTEGER AS n]\n{SCAN}");
assert_eq!(folded(&text), text);
}
#[test]
fn a_cast_of_a_constant_folds_and_one_that_would_raise_does_not() {
let before = format!("Project #1 [CAST('1'::VARCHAR)::INTEGER AS n]\n{SCAN}");
let after = format!("Project #1 [1::INTEGER AS n]\n{SCAN}");
assert_eq!(folded(&before), after);
let raises = format!("Project #1 [CAST('abc'::VARCHAR)::INTEGER AS n]\n{SCAN}");
assert_eq!(folded(&raises), raises);
}
#[test]
fn a_comparison_of_constants_becomes_a_boolean() {
let before = format!("Filter (1::INTEGER < 2::INTEGER)::BOOLEAN\n{SCAN}");
let after = format!("Filter TRUE::BOOLEAN\n{SCAN}");
assert_eq!(folded(&before), after);
}
#[test]
fn a_comparison_against_a_null_is_null_and_the_other_side_goes_with_it() {
let before = format!("Filter (#0.0::INTEGER = NULL::INTEGER)::BOOLEAN\n{SCAN}");
let after = format!("Filter NULL::BOOLEAN\n{SCAN}");
assert_eq!(folded(&before), after);
}
#[test]
fn the_two_comparisons_that_have_an_answer_over_a_null_keep_it() {
let text =
format!("Filter (#0.0::INTEGER IS NOT DISTINCT FROM NULL::INTEGER)::BOOLEAN\n{SCAN}");
assert_eq!(folded(&text), text);
}
#[test]
fn a_true_drops_out_of_an_and_and_a_false_decides_it() {
let before = format!("Filter (TRUE::BOOLEAN AND #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
let after = format!("Filter #0.2::BOOLEAN\n{SCAN}");
assert_eq!(folded(&before), after);
let decided = format!("Filter (FALSE::BOOLEAN AND #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
let all = format!("Filter FALSE::BOOLEAN\n{SCAN}");
assert_eq!(folded(&decided), all);
}
#[test]
fn a_false_drops_out_of_an_or_and_a_true_decides_it() {
let before = format!("Filter (FALSE::BOOLEAN OR #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
let after = format!("Filter #0.2::BOOLEAN\n{SCAN}");
assert_eq!(folded(&before), after);
let decided = format!("Filter (TRUE::BOOLEAN OR #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
let all = format!("Filter TRUE::BOOLEAN\n{SCAN}");
assert_eq!(folded(&decided), all);
}
#[test]
fn a_null_operand_of_an_and_is_kept_because_it_is_neither_the_answer_nor_the_operand() {
let text = format!("Filter (NULL::BOOLEAN AND #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
assert_eq!(folded(&text), text);
}
#[test]
fn a_conjunction_of_constants_is_the_three_valued_answer() {
let before = format!("Filter (NULL::BOOLEAN AND FALSE::BOOLEAN)::BOOLEAN\n{SCAN}");
let after = format!("Filter FALSE::BOOLEAN\n{SCAN}");
assert_eq!(folded(&before), after);
let other = format!("Filter (NULL::BOOLEAN OR TRUE::BOOLEAN)::BOOLEAN\n{SCAN}");
let answer = format!("Filter TRUE::BOOLEAN\n{SCAN}");
assert_eq!(folded(&other), answer);
}
#[test]
fn a_long_conjunction_keeps_the_operands_that_are_not_decided() {
let before = format!(
"Filter (#0.2::BOOLEAN AND TRUE::BOOLEAN AND (#0.0::INTEGER > 1::INTEGER)::BOOLEAN)::BOOLEAN\n{SCAN}"
);
let after = format!(
"Filter (#0.2::BOOLEAN AND (#0.0::INTEGER > 1::INTEGER)::BOOLEAN)::BOOLEAN\n{SCAN}"
);
assert_eq!(folded(&before), after);
}
#[test]
fn an_arm_that_cannot_fire_is_dropped_and_a_null_condition_is_one_of_them() {
let before = format!(
"Project #1 [CASE WHEN FALSE::BOOLEAN THEN 1::INTEGER ELSE #0.0::INTEGER END::INTEGER AS n]\n{SCAN}"
);
let after = format!("Project #1 [#0.0::INTEGER AS n]\n{SCAN}");
assert_eq!(folded(&before), after);
let null = format!(
"Project #1 [CASE WHEN NULL::BOOLEAN THEN 1::INTEGER ELSE #0.0::INTEGER END::INTEGER AS n]\n{SCAN}"
);
assert_eq!(folded(&null), after);
}
#[test]
fn the_first_arm_that_always_fires_cuts_the_ones_after_it() {
let before = format!(
"Project #1 [CASE WHEN #0.2::BOOLEAN THEN 1::INTEGER WHEN TRUE::BOOLEAN THEN 2::INTEGER WHEN #0.2::BOOLEAN THEN 3::INTEGER ELSE 4::INTEGER END::INTEGER AS n]\n{SCAN}"
);
let after = format!(
"Project #1 [CASE WHEN #0.2::BOOLEAN THEN 1::INTEGER ELSE 2::INTEGER END::INTEGER AS n]\n{SCAN}"
);
assert_eq!(folded(&before), after);
}
#[test]
fn a_case_with_no_arm_left_and_no_else_is_null() {
let before = format!(
"Project #1 [CASE WHEN FALSE::BOOLEAN THEN 1::INTEGER END::INTEGER AS n]\n{SCAN}"
);
let after = format!("Project #1 [NULL::INTEGER AS n]\n{SCAN}");
assert_eq!(folded(&before), after);
}
#[test]
fn an_aggregate_keeps_its_place_and_its_arguments_are_folded_under_it() {
let before = format!(
"Aggregate #1 groups=[] aggregates=[sum(\"+\"(1::INTEGER, 2::INTEGER)::INTEGER)::HUGEINT]\n{SCAN}"
);
let after = format!("Aggregate #1 groups=[] aggregates=[sum(3::INTEGER)::HUGEINT]\n{SCAN}");
assert_eq!(folded(&before), after);
}
#[test]
fn a_sort_key_and_a_join_condition_are_folded_too() {
let before =
format!("Sort [\"+\"(1::INTEGER, 1::INTEGER)::INTEGER ASC NULLS LAST]\n{SCAN}");
let after = format!("Sort [2::INTEGER ASC NULLS LAST]\n{SCAN}");
assert_eq!(folded(&before), after);
}
#[test]
fn folding_twice_is_folding_once() {
let before = format!(
"Filter (TRUE::BOOLEAN AND (\"+\"(1::INTEGER, 1::INTEGER)::INTEGER > #0.0::INTEGER)::BOOLEAN)::BOOLEAN\n{SCAN}"
);
let once = folded(&before);
assert_eq!(folded(&once), once);
}
#[test]
fn a_volatile_call_is_not_folded_however_constant_its_arguments_are() {
assert!(VOLATILE.contains(&"random"));
assert!(VOLATILE.contains(&"nextval"));
let text = format!("Project #1 [random()::DOUBLE AS n]\n{SCAN}");
assert_eq!(folded(&text), text);
}
}