1use std::collections::HashMap;
36
37use rudb_common::{LogicalType, Result, Value};
38use rudb_kernels::{Comparison, Connective, call_values, cast_value, combine, compare_values};
39use rudb_plan::{CompareOp, ConjunctionOp, Expr, ExprRef, Node, NodeRef, Plan, Slice, SortKey};
40use rudb_vector::Vector;
41
42use crate::pass::{Context, Pass, top_down};
43use crate::walk;
44
45pub const VOLATILE: [&str; 17] = [
52 "current_connection_id",
53 "current_query",
54 "current_query_id",
55 "current_transaction_id",
56 "currval",
57 "error",
58 "gen_random_uuid",
59 "nextval",
60 "random",
61 "setseed",
62 "setval",
63 "sleep_ms",
64 "stats",
65 "uuid",
66 "uuidv4",
67 "uuidv7",
68 "write_log",
69];
70
71#[derive(Debug, Clone, Copy)]
73pub struct ExpressionRewriter;
74
75impl Pass for ExpressionRewriter {
76 fn name(&self) -> &'static str {
77 "expression_rewriter"
78 }
79
80 fn run(&self, plan: &mut Plan, _context: &Context) -> Result<()> {
81 rewrite(plan);
82 Ok(())
83 }
84}
85
86type Done = HashMap<ExprRef, ExprRef>;
94
95fn rewrite(plan: &mut Plan) {
97 let mut done = Done::new();
98 for node in top_down(plan) {
99 node_expressions(plan, node, &mut done);
100 }
101}
102
103fn node_expressions(plan: &mut Plan, node: NodeRef, done: &mut Done) {
110 match *plan.node(node) {
111 Node::Get { .. }
112 | Node::Dummy
113 | Node::Limit { .. }
114 | Node::SetOp { .. }
115 | Node::CrossProduct { .. } => {}
116 Node::Values { rows, .. } => {
117 let held = plan.row_list(rows).to_vec();
118 let rewritten: Vec<Slice> =
119 held.iter().map(|&row| expr_list(plan, row, done).unwrap_or(row)).collect();
120 if rewritten != held {
121 let rows = plan.add_rows(&rewritten);
122 match plan.node_mut(node) {
123 Node::Values { rows: held, .. } => *held = rows,
124 _ => unreachable!("the node was a values list a moment ago"),
125 }
126 }
127 }
128 Node::TableFunction { args, .. } => {
129 if let Some(rewritten) = expr_list(plan, args, done) {
130 match plan.node_mut(node) {
131 Node::TableFunction { args, .. } => *args = rewritten,
132 _ => unreachable!("the node was a table function a moment ago"),
133 }
134 }
135 }
136 Node::Filter { predicate, .. } => {
137 let rewritten = expression(plan, predicate, done);
138 if rewritten != predicate {
139 match plan.node_mut(node) {
140 Node::Filter { predicate, .. } => *predicate = rewritten,
141 _ => unreachable!("the node was a filter a moment ago"),
142 }
143 }
144 }
145 Node::Project { exprs, .. } => {
146 if let Some(rewritten) = expr_list(plan, exprs, done) {
147 match plan.node_mut(node) {
148 Node::Project { exprs, .. } => *exprs = rewritten,
149 _ => unreachable!("the node was a projection a moment ago"),
150 }
151 }
152 }
153 Node::Aggregate { groups, aggregates, .. } => {
154 let rewritten_groups = expr_list(plan, groups, done);
155 let rewritten_aggregates = expr_list(plan, aggregates, done);
156 match plan.node_mut(node) {
157 Node::Aggregate { groups, aggregates, .. } => {
158 if let Some(rewritten) = rewritten_groups {
159 *groups = rewritten;
160 }
161 if let Some(rewritten) = rewritten_aggregates {
162 *aggregates = rewritten;
163 }
164 }
165 _ => unreachable!("the node was an aggregate a moment ago"),
166 }
167 }
168 Node::Sort { keys, .. } | Node::TopN { keys, .. } => {
169 let held = plan.sort_key_list(keys).to_vec();
170 let rewritten: Vec<SortKey> = held
171 .iter()
172 .map(|key| SortKey { expr: expression(plan, key.expr, done), ..*key })
173 .collect();
174 if rewritten != held {
175 let keys = plan.add_sort_keys(&rewritten);
176 match plan.node_mut(node) {
177 Node::Sort { keys: held, .. } | Node::TopN { keys: held, .. } => *held = keys,
178 _ => unreachable!("the node was a sort a moment ago"),
179 }
180 }
181 }
182 Node::Distinct { on, .. } => {
183 if let Some(rewritten) = expr_list(plan, on, done) {
184 match plan.node_mut(node) {
185 Node::Distinct { on, .. } => *on = rewritten,
186 _ => unreachable!("the node was a distinct a moment ago"),
187 }
188 }
189 }
190 Node::Join { conditions, .. } => {
191 if let Some(rewritten) = expr_list(plan, conditions, done) {
192 match plan.node_mut(node) {
193 Node::Join { conditions, .. } => *conditions = rewritten,
194 _ => unreachable!("the node was a join a moment ago"),
195 }
196 }
197 }
198 }
199}
200
201fn expr_list(plan: &mut Plan, slice: Slice, done: &mut Done) -> Option<Slice> {
203 walk::list(plan, slice, &mut |plan, expr| expression(plan, expr, done))
204}
205
206fn expression(plan: &mut Plan, expr: ExprRef, done: &mut Done) -> ExprRef {
214 if let Some(&already) = done.get(&expr) {
215 return already;
216 }
217 let rebuilt = walk::rebuild(plan, expr, &mut |plan, child| expression(plan, child, done));
218 let simplified = simplify(plan, rebuilt);
219 done.insert(expr, simplified);
220 simplified
221}
222
223fn simplify(plan: &mut Plan, expr: ExprRef) -> ExprRef {
225 if matches!(*plan.expr(expr), Expr::Constant(_)) {
226 return expr;
227 }
228 if let Some(value) = fold(plan, expr) {
229 if let Some(folded) = constant_of(plan, expr, value) {
230 return folded;
231 }
232 }
233 match *plan.expr(expr) {
234 Expr::Conjunction { op, children } => conjunction(plan, expr, op, children),
235 Expr::Case { arms, otherwise } => case(plan, expr, arms, otherwise),
236 Expr::Compare { op, left, right } => null_comparison(plan, expr, op, left, right),
237 _ => expr,
238 }
239}
240
241fn fold(plan: &Plan, expr: ExprRef) -> Option<Value> {
248 match *plan.expr(expr) {
249 Expr::Cast { input, try_cast } => {
250 let inner = constant(plan, input)?;
251 cast_value(&inner, plan.expr_type(expr), try_cast).ok()
252 }
253 Expr::Compare { op, left, right } => {
254 let left = constant(plan, left)?;
255 let right = constant(plan, right)?;
256 compare_values(comparison(op), &left, &right).ok()
257 }
258 Expr::Conjunction { op, children } => {
259 let values = constants(plan, children)?;
260 let vectors: Vec<Vector> = values
261 .into_iter()
262 .map(|value| Vector::constant(LogicalType::Boolean, value, 1))
263 .collect();
264 Some(combine(connective(op), &vectors).ok()?.value_at(0))
265 }
266 Expr::Function { name, args } => {
267 let name = plan.string(name);
268 if VOLATILE.contains(&name) {
269 return None;
270 }
271 let values = constants(plan, args)?;
272 call_values(name, &values, plan.expr_type(expr), None).ok()
277 }
278 _ => None,
279 }
280}
281
282fn constant(plan: &Plan, expr: ExprRef) -> Option<Value> {
284 match *plan.expr(expr) {
285 Expr::Constant(value) => Some(plan.value(value).clone()),
286 _ => None,
287 }
288}
289
290fn constants(plan: &Plan, slice: Slice) -> Option<Vec<Value>> {
292 plan.expr_list(slice).iter().map(|&expr| constant(plan, expr)).collect()
293}
294
295fn constant_of(plan: &mut Plan, expr: ExprRef, value: Value) -> Option<ExprRef> {
301 let ty = plan.expr_type(expr).clone();
302 if !value.is_null() && value.logical_type() != ty {
303 return None;
304 }
305 let held = plan.add_value(value);
306 Some(plan.add_expr(Expr::Constant(held), ty))
307}
308
309fn conjunction(plan: &mut Plan, expr: ExprRef, op: ConjunctionOp, children: Slice) -> ExprRef {
314 let (decides, drops) = match op {
316 ConjunctionOp::And => (false, true),
317 ConjunctionOp::Or => (true, false),
318 };
319 let held = plan.expr_list(children).to_vec();
320 let mut kept = Vec::with_capacity(held.len());
321 for child in held.iter().copied() {
322 match constant(plan, child).as_ref().and_then(Value::as_bool) {
323 Some(known) if known == decides => {
324 return constant_of(plan, expr, Value::Boolean(decides)).unwrap_or(expr);
325 }
326 Some(_) => {}
327 None => kept.push(child),
328 }
329 }
330 if kept.len() == held.len() {
331 return expr;
332 }
333 match kept.as_slice() {
334 [] => constant_of(plan, expr, Value::Boolean(drops)).unwrap_or(expr),
335 [only] => *only,
336 rest => {
337 let children = plan.add_expr_list(rest);
338 plan.add_expr(Expr::Conjunction { op, children }, LogicalType::Boolean)
339 }
340 }
341}
342
343enum Fires {
345 Always,
347 Never,
350 Maybe,
352}
353
354fn fires(plan: &Plan, when: ExprRef) -> Fires {
355 match constant(plan, when) {
356 Some(Value::Boolean(true)) => Fires::Always,
357 Some(Value::Boolean(false) | Value::Null) => Fires::Never,
358 Some(_) | None => Fires::Maybe,
361 }
362}
363
364fn case(plan: &mut Plan, expr: ExprRef, arms: Slice, otherwise: Option<ExprRef>) -> ExprRef {
366 let held = plan.arm_list(arms).to_vec();
367 let mut kept = Vec::with_capacity(held.len());
368 let mut result = otherwise;
369 let mut cut = false;
370 for arm in held.iter().copied() {
371 match fires(plan, arm.when) {
372 Fires::Never => {}
373 Fires::Always => {
374 result = Some(arm.then);
375 cut = true;
376 break;
377 }
378 Fires::Maybe => kept.push(arm),
379 }
380 }
381 if kept.len() == held.len() && !cut {
382 return expr;
383 }
384 if kept.is_empty() {
385 return match result {
389 Some(only) if plan.expr_type(only) == plan.expr_type(expr) => only,
390 Some(_) => expr,
391 None => constant_of(plan, expr, Value::Null).unwrap_or(expr),
392 };
393 }
394 let ty = plan.expr_type(expr).clone();
395 let arms = plan.add_arms(&kept);
396 plan.add_expr(Expr::Case { arms, otherwise: result }, ty)
397}
398
399fn null_comparison(
408 plan: &mut Plan,
409 expr: ExprRef,
410 op: CompareOp,
411 left: ExprRef,
412 right: ExprRef,
413) -> ExprRef {
414 if matches!(op, CompareOp::DistinctFrom | CompareOp::NotDistinctFrom) {
415 return expr;
416 }
417 let is_null = |side| constant(plan, side).is_some_and(|value| value.is_null());
418 if is_null(left) || is_null(right) {
419 constant_of(plan, expr, Value::Null).unwrap_or(expr)
420 } else {
421 expr
422 }
423}
424
425fn comparison(op: CompareOp) -> Comparison {
434 match op {
435 CompareOp::Equal => Comparison::Equal,
436 CompareOp::NotEqual => Comparison::NotEqual,
437 CompareOp::Less => Comparison::Less,
438 CompareOp::LessOrEqual => Comparison::LessOrEqual,
439 CompareOp::Greater => Comparison::Greater,
440 CompareOp::GreaterOrEqual => Comparison::GreaterOrEqual,
441 CompareOp::DistinctFrom => Comparison::DistinctFrom,
442 CompareOp::NotDistinctFrom => Comparison::NotDistinctFrom,
443 }
444}
445
446fn connective(op: ConjunctionOp) -> Connective {
448 match op {
449 ConjunctionOp::And => Connective::And,
450 ConjunctionOp::Or => Connective::Or,
451 }
452}
453
454#[cfg(test)]
455mod tests {
456 use super::{ExpressionRewriter, VOLATILE};
457 use crate::pass::{Context, Pass};
458 use rudb_plan::Plan;
459
460 fn folded(text: &str) -> String {
462 let mut plan =
463 Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
464 ExpressionRewriter
465 .run(&mut plan, &Context::new())
466 .unwrap_or_else(|error| panic!("{text} did not fold: {error}"));
467 plan.validate().unwrap_or_else(|error| panic!("{text} folded to a bad plan: {error}"));
468 plan.to_string()
469 }
470
471 const SCAN: &str = " Get memory.main.t AS t #0 [a::INTEGER, b::VARCHAR, c::BOOLEAN]\n";
472
473 #[test]
474 fn arithmetic_over_constants_becomes_the_number() {
475 let before = format!("Project #1 [\"+\"(2::INTEGER, 3::INTEGER)::INTEGER AS n]\n{SCAN}");
476 let after = format!("Project #1 [5::INTEGER AS n]\n{SCAN}");
477 assert_eq!(folded(&before), after);
478 }
479
480 #[test]
481 fn a_nest_of_constants_folds_all_the_way_up_in_one_walk() {
482 let before = format!(
485 "Project #1 [\"+\"(\"+\"(1::INTEGER, 2::INTEGER)::INTEGER, 3::INTEGER)::INTEGER AS n]\n{SCAN}"
486 );
487 let after = format!("Project #1 [6::INTEGER AS n]\n{SCAN}");
488 assert_eq!(folded(&before), after);
489 }
490
491 #[test]
492 fn a_call_with_a_column_in_it_is_left_alone() {
493 let text = format!("Project #1 [\"+\"(#0.0::INTEGER, 3::INTEGER)::INTEGER AS n]\n{SCAN}");
494 assert_eq!(folded(&text), text);
495 }
496
497 #[test]
498 fn a_cast_of_a_constant_folds_and_one_that_would_raise_does_not() {
499 let before = format!("Project #1 [CAST('1'::VARCHAR)::INTEGER AS n]\n{SCAN}");
500 let after = format!("Project #1 [1::INTEGER AS n]\n{SCAN}");
501 assert_eq!(folded(&before), after);
502 let raises = format!("Project #1 [CAST('abc'::VARCHAR)::INTEGER AS n]\n{SCAN}");
505 assert_eq!(folded(&raises), raises);
506 }
507
508 #[test]
509 fn a_comparison_of_constants_becomes_a_boolean() {
510 let before = format!("Filter (1::INTEGER < 2::INTEGER)::BOOLEAN\n{SCAN}");
511 let after = format!("Filter TRUE::BOOLEAN\n{SCAN}");
512 assert_eq!(folded(&before), after);
513 }
514
515 #[test]
516 fn a_comparison_against_a_null_is_null_and_the_other_side_goes_with_it() {
517 let before = format!("Filter (#0.0::INTEGER = NULL::INTEGER)::BOOLEAN\n{SCAN}");
518 let after = format!("Filter NULL::BOOLEAN\n{SCAN}");
519 assert_eq!(folded(&before), after);
520 }
521
522 #[test]
523 fn the_two_comparisons_that_have_an_answer_over_a_null_keep_it() {
524 let text =
527 format!("Filter (#0.0::INTEGER IS NOT DISTINCT FROM NULL::INTEGER)::BOOLEAN\n{SCAN}");
528 assert_eq!(folded(&text), text);
529 }
530
531 #[test]
532 fn a_true_drops_out_of_an_and_and_a_false_decides_it() {
533 let before = format!("Filter (TRUE::BOOLEAN AND #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
534 let after = format!("Filter #0.2::BOOLEAN\n{SCAN}");
535 assert_eq!(folded(&before), after);
536 let decided = format!("Filter (FALSE::BOOLEAN AND #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
537 let all = format!("Filter FALSE::BOOLEAN\n{SCAN}");
538 assert_eq!(folded(&decided), all);
539 }
540
541 #[test]
542 fn a_false_drops_out_of_an_or_and_a_true_decides_it() {
543 let before = format!("Filter (FALSE::BOOLEAN OR #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
544 let after = format!("Filter #0.2::BOOLEAN\n{SCAN}");
545 assert_eq!(folded(&before), after);
546 let decided = format!("Filter (TRUE::BOOLEAN OR #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
547 let all = format!("Filter TRUE::BOOLEAN\n{SCAN}");
548 assert_eq!(folded(&decided), all);
549 }
550
551 #[test]
552 fn a_null_operand_of_an_and_is_kept_because_it_is_neither_the_answer_nor_the_operand() {
553 let text = format!("Filter (NULL::BOOLEAN AND #0.2::BOOLEAN)::BOOLEAN\n{SCAN}");
554 assert_eq!(folded(&text), text);
555 }
556
557 #[test]
558 fn a_conjunction_of_constants_is_the_three_valued_answer() {
559 let before = format!("Filter (NULL::BOOLEAN AND FALSE::BOOLEAN)::BOOLEAN\n{SCAN}");
562 let after = format!("Filter FALSE::BOOLEAN\n{SCAN}");
563 assert_eq!(folded(&before), after);
564 let other = format!("Filter (NULL::BOOLEAN OR TRUE::BOOLEAN)::BOOLEAN\n{SCAN}");
565 let answer = format!("Filter TRUE::BOOLEAN\n{SCAN}");
566 assert_eq!(folded(&other), answer);
567 }
568
569 #[test]
570 fn a_long_conjunction_keeps_the_operands_that_are_not_decided() {
571 let before = format!(
572 "Filter (#0.2::BOOLEAN AND TRUE::BOOLEAN AND (#0.0::INTEGER > 1::INTEGER)::BOOLEAN)::BOOLEAN\n{SCAN}"
573 );
574 let after = format!(
575 "Filter (#0.2::BOOLEAN AND (#0.0::INTEGER > 1::INTEGER)::BOOLEAN)::BOOLEAN\n{SCAN}"
576 );
577 assert_eq!(folded(&before), after);
578 }
579
580 #[test]
581 fn an_arm_that_cannot_fire_is_dropped_and_a_null_condition_is_one_of_them() {
582 let before = format!(
583 "Project #1 [CASE WHEN FALSE::BOOLEAN THEN 1::INTEGER ELSE #0.0::INTEGER END::INTEGER AS n]\n{SCAN}"
584 );
585 let after = format!("Project #1 [#0.0::INTEGER AS n]\n{SCAN}");
586 assert_eq!(folded(&before), after);
587 let null = format!(
588 "Project #1 [CASE WHEN NULL::BOOLEAN THEN 1::INTEGER ELSE #0.0::INTEGER END::INTEGER AS n]\n{SCAN}"
589 );
590 assert_eq!(folded(&null), after);
591 }
592
593 #[test]
594 fn the_first_arm_that_always_fires_cuts_the_ones_after_it() {
595 let before = format!(
596 "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}"
597 );
598 let after = format!(
599 "Project #1 [CASE WHEN #0.2::BOOLEAN THEN 1::INTEGER ELSE 2::INTEGER END::INTEGER AS n]\n{SCAN}"
600 );
601 assert_eq!(folded(&before), after);
602 }
603
604 #[test]
605 fn a_case_with_no_arm_left_and_no_else_is_null() {
606 let before = format!(
607 "Project #1 [CASE WHEN FALSE::BOOLEAN THEN 1::INTEGER END::INTEGER AS n]\n{SCAN}"
608 );
609 let after = format!("Project #1 [NULL::INTEGER AS n]\n{SCAN}");
610 assert_eq!(folded(&before), after);
611 }
612
613 #[test]
614 fn an_aggregate_keeps_its_place_and_its_arguments_are_folded_under_it() {
615 let before = format!(
618 "Aggregate #1 groups=[] aggregates=[sum(\"+\"(1::INTEGER, 2::INTEGER)::INTEGER)::HUGEINT]\n{SCAN}"
619 );
620 let after = format!("Aggregate #1 groups=[] aggregates=[sum(3::INTEGER)::HUGEINT]\n{SCAN}");
621 assert_eq!(folded(&before), after);
622 }
623
624 #[test]
625 fn a_sort_key_and_a_join_condition_are_folded_too() {
626 let before =
627 format!("Sort [\"+\"(1::INTEGER, 1::INTEGER)::INTEGER ASC NULLS LAST]\n{SCAN}");
628 let after = format!("Sort [2::INTEGER ASC NULLS LAST]\n{SCAN}");
629 assert_eq!(folded(&before), after);
630 }
631
632 #[test]
633 fn folding_twice_is_folding_once() {
634 let before = format!(
635 "Filter (TRUE::BOOLEAN AND (\"+\"(1::INTEGER, 1::INTEGER)::INTEGER > #0.0::INTEGER)::BOOLEAN)::BOOLEAN\n{SCAN}"
636 );
637 let once = folded(&before);
638 assert_eq!(folded(&once), once);
639 }
640
641 #[test]
642 fn a_volatile_call_is_not_folded_however_constant_its_arguments_are() {
643 assert!(VOLATILE.contains(&"random"));
646 assert!(VOLATILE.contains(&"nextval"));
647 let text = format!("Project #1 [random()::DOUBLE AS n]\n{SCAN}");
648 assert_eq!(folded(&text), text);
649 }
650}