1use crate::sqlselect::Expr;
85use serde_json::Value;
86use std::collections::HashMap;
87
88const PURE_FUNCS: &[&str] = &[
94 "lower", "upper", "length", "char_length", "character_length", "coalesce",
95 "nullif", "int2", "int4", "int8", "text", "quote_ident", "format_type",
96 "array_to_string", "current_schema", "current_database", "current_catalog",
97 "current_user", "session_user", "user", "version", "pg_get_userbyid",
98 "pg_table_is_visible", "pg_type_is_visible", "pg_function_is_visible",
99 "pg_encoding_to_char", "pg_get_expr", "pg_get_indexdef",
100 "pg_get_constraintdef",
101];
102
103#[derive(Debug, Clone, Default)]
105pub struct Pushdown {
106 pub per_binding: HashMap<String, Vec<Expr>>,
108 pub refusals: Vec<String>,
110}
111
112impl Pushdown {
113 pub fn for_binding(&self, binding: &str) -> Option<&Vec<Expr>> {
114 self.per_binding.get(&binding.to_ascii_lowercase())
115 }
116
117 pub fn pushed_count(&self) -> usize {
118 self.per_binding.values().map(|v| v.len()).sum()
119 }
120}
121
122fn conjuncts<'a>(e: &'a Expr, out: &mut Vec<&'a Expr>) {
127 match e {
128 Expr::Binary { op, left, right } if op == "AND" => {
129 conjuncts(left, out);
130 conjuncts(right, out);
131 }
132 other => out.push(other),
133 }
134}
135
136enum Reads {
138 One(String),
140 Constant,
143 Refused(&'static str),
144}
145
146fn reads(e: &Expr, known: &[String]) -> Reads {
147 let mut seen: Vec<String> = vec![];
148 let mut why: Option<&'static str> = None;
149 walk(e, known, &mut seen, &mut why);
150 if let Some(w) = why {
151 return Reads::Refused(w);
152 }
153 match seen.len() {
154 0 => Reads::Constant,
155 1 => Reads::One(seen.pop().expect("one")),
156 _ => Reads::Refused("spans more than one relation"),
157 }
158}
159
160fn walk(e: &Expr, known: &[String], seen: &mut Vec<String>, why: &mut Option<&'static str>) {
161 match e {
162 Expr::Column { qual, .. } => match qual {
163 Some(q) => {
164 let lower = q.to_ascii_lowercase();
165 if !known.iter().any(|b| b.eq_ignore_ascii_case(q)) {
166 *why = Some("references an unknown relation");
169 } else if !seen.contains(&lower) {
170 seen.push(lower);
171 }
172 }
173 None => *why = Some("unqualified column cannot be attributed to a relation"),
177 },
178 Expr::Literal(_) => {}
179 Expr::Star | Expr::QualifiedStar(_) => *why = Some("contains `*`"),
180 Expr::Func { name, args } => {
181 if !PURE_FUNCS.iter().any(|f| f.eq_ignore_ascii_case(name)) {
182 *why = Some("calls a function not known to be pure");
183 }
184 for a in args {
185 walk(a, known, seen, why);
186 }
187 }
188 Expr::Agg { .. } => *why = Some("contains an aggregate"),
192 Expr::Case { operand, whens, else_ } => {
193 if let Some(o) = operand {
194 walk(o, known, seen, why);
195 }
196 for (w, t) in whens {
197 walk(w, known, seen, why);
198 walk(t, known, seen, why);
199 }
200 if let Some(x) = else_ {
201 walk(x, known, seen, why);
202 }
203 }
204 Expr::Binary { left, right, .. } => {
205 walk(left, known, seen, why);
206 walk(right, known, seen, why);
207 }
208 Expr::Unary { expr, .. } | Expr::Cast { expr, .. } | Expr::IsNull { expr, .. } => {
209 walk(expr, known, seen, why)
210 }
211 Expr::InList { expr, list, .. } => {
212 walk(expr, known, seen, why);
213 for i in list {
214 walk(i, known, seen, why);
215 }
216 }
217 Expr::Index { expr, index } => {
218 walk(expr, known, seen, why);
219 walk(index, known, seen, why);
220 }
221 Expr::ArrayLit(items) => {
222 for i in items {
223 walk(i, known, seen, why);
224 }
225 }
226 Expr::Subquery(_) | Expr::Exists { .. } | Expr::ArrayQuery(_) | Expr::InSubquery { .. } => {
230 *why = Some("contains a subquery")
231 }
232 Expr::Quantified { left, right, .. } => {
233 walk(left, known, seen, why);
234 walk(right, known, seen, why);
235 }
236 }
237}
238
239pub fn nullable_bindings(sel: &crate::sqlselect::Select) -> Vec<String> {
245 use crate::sqlselect::JoinKind;
246 let mut out: Vec<String> = vec![];
247 let mut accumulated: Vec<String> = sel
248 .from
249 .iter()
250 .map(|t| t.binding().to_ascii_lowercase())
251 .collect();
252
253 for j in &sel.joins {
254 let rb = j.table.binding().to_ascii_lowercase();
255 if matches!(j.kind, JoinKind::Left | JoinKind::Full) && !out.contains(&rb) {
258 out.push(rb.clone());
259 }
260 if matches!(j.kind, JoinKind::Right | JoinKind::Full) {
264 for a in &accumulated {
265 if !out.contains(a) {
266 out.push(a.clone());
267 }
268 }
269 }
270 accumulated.push(rb);
271 }
272 out
273}
274
275pub fn plan(
281 where_: Option<&Expr>,
282 bindings: &[String],
283 nullable: &[String],
284) -> Pushdown {
285 let mut out = Pushdown::default();
286 let Some(w) = where_ else { return out };
287
288 if bindings.len() < 2 {
291 return out;
292 }
293
294 let mut parts = vec![];
295 conjuncts(w, &mut parts);
296 for p in parts {
297 match reads(p, bindings) {
298 Reads::One(b) if nullable.iter().any(|n| n.eq_ignore_ascii_case(&b)) => {
299 out.refusals.push(format!(
302 "Filter retained above join: predicate references nullable \
303 side of an outer join ({b})"
304 ));
305 }
306 Reads::One(b) => out.per_binding.entry(b).or_default().push(p.clone()),
307 Reads::Constant => out
308 .refusals
309 .push("Filter retained above join: predicate reads no column".into()),
310 Reads::Refused(why) => out
311 .refusals
312 .push(format!("Filter retained above join: {why}")),
313 }
314 }
315 out
316}
317
318pub fn to_nql_predicate(e: &Expr, binding: &str, strict_qual: bool) -> Option<String> {
366 match e {
367 Expr::Column { qual, name } => match qual {
368 Some(q) if q.eq_ignore_ascii_case(binding) => Some(name.clone()),
369 Some(_) => None,
370 None if strict_qual => None,
371 None => Some(name.clone()),
372 },
373 Expr::Literal(v) => nql_literal(v),
374 Expr::Binary { op, left, right } => {
375 let o = op.to_ascii_uppercase();
376 let l = to_nql_predicate(left, binding, strict_qual)?;
377 let r = to_nql_predicate(right, binding, strict_qual)?;
378 match o.as_str() {
379 "=" | "!=" | ">" | "<" | ">=" | "<=" | "LIKE" => Some(format!("{} {} {}", l, o, r)),
381 "<>" => Some(format!("{} != {}", l, r)),
382 "AND" => Some(format!("({} AND {})", l, r)),
385 "OR" => Some(format!("({} OR {})", l, r)),
386 _ => None,
387 }
388 }
389 Expr::InList { expr, list, negated: false } => {
390 let l = to_nql_predicate(expr, binding, strict_qual)?;
391 let mut items = Vec::with_capacity(list.len());
392 for it in list {
393 items.push(to_nql_predicate(it, binding, strict_qual)?);
394 }
395 if items.is_empty() {
396 return None;
397 }
398 Some(format!("{} IN ({})", l, items.join(", ")))
399 }
400 _ => None,
402 }
403}
404
405fn nql_literal(v: &Value) -> Option<String> {
407 match v {
408 Value::String(s) => Some(format!("\"{}\"", s.replace('\\', "\\\\").replace('"', "\\\""))),
410 Value::Number(n) => Some(n.to_string()),
411 Value::Bool(b) => Some(if *b { "TRUE".into() } else { "FALSE".into() }),
412 _ => None,
415 }
416}
417
418pub fn nql_prefilter(
426 where_: Option<&Expr>,
427 binding: &str,
428 bindings: &[String],
429 nullable: &[String],
430) -> Option<String> {
431 if nullable.iter().any(|n| n.eq_ignore_ascii_case(binding)) {
435 return None;
436 }
437 let w = where_?;
438 let strict = bindings.len() > 1;
439 let mut parts = vec![];
440 conjuncts(w, &mut parts);
441 let kept: Vec<String> = parts
442 .iter()
443 .filter_map(|p| to_nql_predicate(p, binding, strict))
444 .collect();
445 if kept.is_empty() {
446 None
447 } else {
448 Some(kept.join(" AND "))
449 }
450}
451
452#[cfg(test)]
453mod tests {
454 use super::*;
455 use crate::sqlselect::parse;
456
457 fn plan_for(sql: &str) -> Pushdown {
458 let sel = parse(sql).expect("parses");
459 let mut b = vec![];
460 if let Some(f) = &sel.from {
461 b.push(f.binding());
462 }
463 for j in &sel.joins {
464 b.push(j.table.binding());
465 }
466 let nullable = nullable_bindings(&sel);
467 plan(sel.where_.as_ref(), &b, &nullable)
468 }
469
470 #[test]
471 fn a_single_relation_predicate_is_pushed_to_that_relation() {
472 let p = plan_for("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE a.v > 5");
473 assert_eq!(p.pushed_count(), 1);
474 assert_eq!(p.for_binding("a").map(|v| v.len()), Some(1));
475 assert!(p.for_binding("b").is_none());
476 assert!(p.refusals.is_empty(), "{:?}", p.refusals);
477 }
478
479 #[test]
480 fn conjuncts_are_pushed_to_their_own_relations_independently() {
481 let p = plan_for(
482 "SELECT 1 FROM a JOIN b ON a.x = b.x WHERE a.v > 5 AND b.w < 2 AND a.z = 'q'",
483 );
484 assert_eq!(p.pushed_count(), 3);
485 assert_eq!(p.for_binding("a").map(|v| v.len()), Some(2));
486 assert_eq!(p.for_binding("b").map(|v| v.len()), Some(1));
487 }
488
489 #[test]
490 fn a_predicate_on_the_nullable_side_of_a_left_join_is_REFUSED() {
491 let p = plan_for("SELECT 1 FROM a LEFT JOIN b ON a.x = b.x WHERE b.w = 5");
497 assert_eq!(p.pushed_count(), 0);
498 assert!(p.refusals[0].contains("nullable side"), "{:?}", p.refusals);
499 }
500
501 #[test]
502 fn the_non_nullable_side_of_a_left_join_is_still_pushed() {
503 let p = plan_for("SELECT 1 FROM a LEFT JOIN b ON a.x = b.x WHERE a.v > 5");
507 assert_eq!(p.for_binding("a").map(|v| v.len()), Some(1));
508 assert!(p.refusals.is_empty(), "{:?}", p.refusals);
509 }
510
511 #[test]
512 fn a_right_join_makes_the_LEFT_side_nullable_including_the_from_relation() {
513 let p = plan_for("SELECT 1 FROM a RIGHT JOIN b ON a.x = b.x WHERE a.v > 5");
514 assert_eq!(p.pushed_count(), 0, "a is synthesised by the RIGHT join");
515 assert!(p.refusals[0].contains("nullable side"), "{:?}", p.refusals);
516 let p = plan_for("SELECT 1 FROM a RIGHT JOIN b ON a.x = b.x WHERE b.w > 5");
518 assert_eq!(p.for_binding("b").map(|v| v.len()), Some(1));
519 }
520
521 #[test]
522 fn a_full_join_makes_both_sides_nullable() {
523 for w in ["a.v > 5", "b.w > 5"] {
524 let p = plan_for(&format!("SELECT 1 FROM a FULL JOIN b ON a.x = b.x WHERE {w}"));
525 assert_eq!(p.pushed_count(), 0, "{w}");
526 }
527 }
528
529 #[test]
530 fn a_later_right_join_retroactively_protects_earlier_relations() {
531 let sel = parse(
536 "SELECT 1 FROM a JOIN b ON a.x = b.x RIGHT JOIN c ON b.y = c.y \
537 WHERE a.v > 1 AND b.w > 1 AND c.z > 1",
538 )
539 .expect("parses");
540 let nullable = nullable_bindings(&sel);
541 assert!(nullable.contains(&"a".to_string()), "{nullable:?}");
542 assert!(nullable.contains(&"b".to_string()), "{nullable:?}");
543 assert!(!nullable.contains(&"c".to_string()), "c is never synthesised");
544
545 let p = plan_for(
546 "SELECT 1 FROM a JOIN b ON a.x = b.x RIGHT JOIN c ON b.y = c.y \
547 WHERE a.v > 1 AND b.w > 1 AND c.z > 1",
548 );
549 assert_eq!(p.pushed_count(), 1, "only c");
550 assert_eq!(p.for_binding("c").map(|v| v.len()), Some(1));
551 assert_eq!(p.refusals.len(), 2);
552 }
553
554 #[test]
555 fn an_all_inner_query_can_push_everything() {
556 let p = plan_for(
557 "SELECT 1 FROM a JOIN b ON a.x = b.x JOIN c ON b.y = c.y \
558 WHERE a.v > 1 AND b.w > 1 AND c.z > 1",
559 );
560 assert_eq!(p.pushed_count(), 3);
561 assert!(p.refusals.is_empty());
562 assert!(nullable_bindings(&parse(
563 "SELECT 1 FROM a JOIN b ON a.x = b.x JOIN c ON b.y = c.y"
564 ).unwrap()).is_empty());
565 }
566
567 #[test]
568 fn a_predicate_spanning_two_relations_is_refused_with_a_reason() {
569 let p = plan_for("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE a.v > b.w");
570 assert_eq!(p.pushed_count(), 0);
571 assert_eq!(p.refusals.len(), 1);
572 assert!(p.refusals[0].contains("spans more than one relation"), "{:?}", p.refusals);
573 }
574
575 #[test]
576 fn or_is_never_split() {
577 let p = plan_for("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE a.v > 5 OR b.w < 2");
580 assert_eq!(p.pushed_count(), 0);
581 assert_eq!(p.refusals.len(), 1);
582 }
583
584 #[test]
585 fn an_or_of_one_relation_is_also_refused_today() {
586 let p = plan_for("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE a.v > 5 OR a.v < 1");
590 assert_eq!(p.pushed_count(), 1, "one conjunct, one relation");
591 }
592
593 #[test]
594 fn an_unqualified_column_is_refused() {
595 let p = plan_for("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE v > 5");
596 assert_eq!(p.pushed_count(), 0);
597 assert!(p.refusals[0].contains("unqualified"), "{:?}", p.refusals);
598 }
599
600 #[test]
601 fn a_constant_predicate_is_refused_as_pointless() {
602 let p = plan_for("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE 1 = 1");
603 assert_eq!(p.pushed_count(), 0);
604 assert!(p.refusals[0].contains("reads no column"), "{:?}", p.refusals);
605 }
606
607 #[test]
608 fn a_volatile_function_is_refused_because_the_allowlist_is_fail_safe() {
609 let sel = parse("SELECT 1 FROM a JOIN b ON a.x = b.x").expect("parses");
610 let _ = sel;
611 let pred = Expr::Binary {
612 op: "=".into(),
613 left: Box::new(Expr::Func {
614 name: "random".into(),
615 args: vec![Expr::Column { qual: Some("a".into()), name: "v".into() }],
616 }),
617 right: Box::new(Expr::Literal(serde_json::json!(1))),
618 };
619 let p = plan(Some(&pred), &["a".into(), "b".into()], &[]);
620 assert_eq!(p.pushed_count(), 0);
621 assert!(p.refusals[0].contains("not known to be pure"), "{:?}", p.refusals);
622 }
623
624 #[test]
625 fn pure_functions_and_postfix_operators_are_pushable() {
626 for (w, want) in [
630 ("lower(a.name) = 'x'", 1),
631 ("a.v IS NULL", 1),
632 ("a.v IS NOT NULL", 1),
633 ("a.v IN (1, 2, 3)", 1),
634 ("a.v NOT IN (1, 2)", 1),
635 ("a.v BETWEEN 1 AND 9", 2),
636 ("a.v NOT BETWEEN 1 AND 9", 1),
637 ("coalesce(a.v, 0) > 1", 1),
638 ("a.v::text = '5'", 1),
639 ("NOT (a.v = 3)", 1),
640 ("CASE WHEN a.v > 1 THEN true ELSE false END", 1),
641 ] {
642 let p = plan_for(&format!("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE {w}"));
643 assert_eq!(
644 p.pushed_count(), want,
645 "{w} should push {want}: {:?}", p.refusals
646 );
647 assert!(p.refusals.is_empty(), "{w}: {:?}", p.refusals);
648 }
649 }
650
651 #[test]
652 fn nothing_is_pushed_without_a_join_because_there_is_nothing_to_push_below() {
653 let p = plan_for("SELECT 1 FROM a WHERE a.v > 5");
654 assert_eq!(p.pushed_count(), 0);
655 assert!(p.refusals.is_empty());
657 }
658
659 #[test]
660 fn an_unknown_relation_is_left_to_the_evaluator_to_report() {
661 let p = plan_for("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE zz.v > 5");
662 assert_eq!(p.pushed_count(), 0);
663 assert!(p.refusals[0].contains("unknown relation"), "{:?}", p.refusals);
664 }
665
666 #[test]
667 fn a_binding_is_matched_case_insensitively() {
668 let p = plan_for("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE A.v > 5");
671 assert_eq!(p.for_binding("a").map(|v| v.len()), Some(1));
672 assert_eq!(p.for_binding("A").map(|v| v.len()), Some(1));
673 }
674
675 fn pre(sql: &str, binding: &str) -> Option<String> {
679 let sel = parse(sql).expect("parses");
680 let bindings: Vec<String> = sel
681 .from
682 .iter()
683 .map(|t| t.binding())
684 .chain(sel.joins.iter().map(|j| j.table.binding()))
685 .collect();
686 let nullable = super::nullable_bindings(&sel);
687 super::nql_prefilter(sel.where_.as_ref(), binding, &bindings, &nullable)
688 }
689
690 #[test]
691 fn the_predicate_reaches_the_scan_in_nqls_spelling() {
692 assert_eq!(pre("SELECT 1 FROM orders WHERE status = 'paid'", "orders").as_deref(),
693 Some("status = \"paid\""));
694 assert_eq!(pre("SELECT 1 FROM orders WHERE total <> 5", "orders").as_deref(),
696 Some("total != 5"));
697 assert_eq!(pre("SELECT 1 FROM orders WHERE total >= 100", "orders").as_deref(),
698 Some("total >= 100"));
699 assert_eq!(pre("SELECT 1 FROM orders WHERE status LIKE 'pa%'", "orders").as_deref(),
700 Some("status LIKE \"pa%\""));
701 assert_eq!(pre("SELECT 1 FROM orders WHERE status IN ('paid','open')", "orders").as_deref(),
702 Some("status IN (\"paid\", \"open\")"));
703 assert_eq!(pre("SELECT 1 FROM orders WHERE a = 1 OR b = 2", "orders").as_deref(),
704 Some("(a = 1 OR b = 2)"));
705 assert_eq!(pre("SELECT 1 FROM orders WHERE s = 'a\"b'", "orders").as_deref(),
707 Some("s = \"a\\\"b\""));
708 }
709
710 #[test]
714 fn anything_that_could_drop_a_row_sql_keeps_is_refused() {
715 for sql in [
716 "SELECT 1 FROM orders WHERE NOT (status = 'paid')",
718 "SELECT 1 FROM orders WHERE status IS NULL",
719 "SELECT 1 FROM orders WHERE status IS NOT NULL",
720 "SELECT 1 FROM orders WHERE status NOT IN ('paid')",
721 "SELECT 1 FROM orders WHERE status = NULL",
722 "SELECT 1 FROM orders WHERE total + 1 > 5",
724 "SELECT 1 FROM orders WHERE lower(status) = 'paid'",
725 "SELECT 1 FROM orders WHERE total::text = '5'",
726 ] {
727 assert_eq!(pre(sql, "orders"), None, "{}", sql);
728 }
729 }
730
731 #[test]
732 fn a_conjunction_pushes_the_part_it_can_and_keeps_the_rest_above() {
733 assert_eq!(pre("SELECT 1 FROM orders WHERE status = 'paid' AND lower(x) = 'y'", "orders")
737 .as_deref(),
738 Some("status = \"paid\""));
739 assert_eq!(pre("SELECT 1 FROM orders WHERE status = 'paid' OR lower(x) = 'y'", "orders"),
742 None);
743 }
744
745 #[test]
746 fn another_relations_predicate_never_reaches_this_scan() {
747 let sql = "SELECT 1 FROM orders o JOIN drivers d ON o.driver = d._id \
748 WHERE o.status = 'paid' AND d.name = 'Bob'";
749 assert_eq!(pre(sql, "o").as_deref(), Some("status = \"paid\""));
750 assert_eq!(pre(sql, "d").as_deref(), Some("name = \"Bob\""));
751 assert_eq!(pre("SELECT 1 FROM a JOIN b ON a.x = b.x WHERE v > 5", "a"), None);
754 assert_eq!(pre("SELECT 1 FROM orders WHERE v > 5", "orders").as_deref(), Some("v > 5"));
756 }
757
758 #[test]
759 fn the_nullable_side_of_an_outer_join_is_never_pre_filtered() {
760 let sql = "SELECT 1 FROM a LEFT JOIN b ON a.x = b.x WHERE b.v = 5";
763 assert_eq!(pre(sql, "b"), None);
764 }
765}