reach.rsannotatedreach.rssource221 lines · 8.8 KB · raw
1//! Reach: the condition under which Postgres evaluates each call.
2//!
3//! The scan answers every call before its quals and projection run, but
4//! Postgres short-circuits: an `OR` arm runs only when the earlier arms
5//! are not true, an `AND` arm only when they are not false, a `CASE`
6//! result only when its `WHEN` is true and the earlier ones are not, a
7//! later qual only when the earlier ones are true, and the projection
8//! only for a row that passed the quals. A call it never reaches must
9//! not be sent, since its answer would be paid for and never read.
10//!
11//! This runs in the plan post-pass, after setrefs, over each jev scan's
12//! final qual and targetlist, where a call is an `INDEX_VAR` column past
13//! the child's. Each call's condition is the `OR`, over its occurrences,
14//! of the `AND` of the conditions on the path to it, and it goes in
15//! `custom_exprs` after the arguments (`plan_custom_path` leaves a
16//! `true` there). The scan evaluates it on each child row and leaves an
17//! unreached call NULL, which Postgres then never reads.
18//!
19//! A condition Postgres could not evaluate a second time with the same
20//! result is dropped, which only widens the reach: one that reads a call's
21//! answer (it is not known yet), a volatile function, a SubPlan, or a
22//! `CASE x WHEN`'s placeholder. Dropped, the call is judged as before.
23
24use pgrx::pg_sys;
25
26use super::ffi::{append, ints, is_a, len, ptrs, walk};
27use super::plan::tables;
28
29/// Sets the reach of every call of every jev scan in `plan`.
30///
31/// # Safety
32/// `plan` is NULL or a finished plan tree, after setrefs.
33pub(super) unsafe fn plan(plan: *mut pg_sys::Plan, children: unsafe fn(*mut pg_sys::Plan) -> Vec<*mut pg_sys::Plan>) {
34    unsafe {
35        if plan.is_null() {
36            return;
37        }
38        if is_a(plan.cast(), pg_sys::NodeTag::T_CustomScan) {
39            let scan = plan.cast::<pg_sys::CustomScan>();
40            if std::ptr::eq((*scan).methods, &tables().scan) {
41                set(scan);
42            }
43        }
44        for child in children(plan) {
45            self::plan(child, children);
46        }
47    }
48}
49
50/// Replaces the trailing `true` per call in `scan`'s `custom_exprs`.
51unsafe fn set(scan: *mut pg_sys::CustomScan) {
52    unsafe {
53        let mut private = ints((*scan).custom_private);
54        let prefix = private.next().expect("the child's column count") as i16;
55        let calls = len((*scan).custom_scan_tlist) - prefix as usize;
56        let exprs: Vec<*mut pg_sys::Node> = ptrs((*scan).custom_exprs).collect();
57        assert!(exprs.len() >= calls, "a reach per call");
58
59        let mut found = Found { prefix, occurrences: vec![Vec::new(); calls] };
60        // ExecScan runs the quals in order, stopping at the first that is
61        // not true, then projects the rows that passed.
62        let mut guards = Vec::new();
63        for qual in ptrs::<pg_sys::Node>((*scan).scan.plan.qual) {
64            found.visit(qual, &guards);
65            guards.push(test(qual, pg_sys::BoolTestType::IS_TRUE));
66        }
67        found.visit((*scan).scan.plan.targetlist.cast(), &guards);
68
69        let head = exprs.len() - calls;
70        let mut out = std::ptr::null_mut();
71        out = append(out, exprs[..head].iter().copied());
72        for occurrences in found.occurrences {
73            out = append(out, [reach(prefix, occurrences)]);
74        }
75        (*scan).custom_exprs = out;
76    }
77}
78
79/// `true` when a call is unused or reached unconditionally somewhere,
80/// else the `OR` of its occurrences' usable conditions.
81unsafe fn reach(prefix: i16, occurrences: Vec<Vec<*mut pg_sys::Node>>) -> *mut pg_sys::Node {
82    unsafe {
83        let mut arms = Vec::new();
84        for guards in occurrences {
85            let usable: Vec<_> = guards.into_iter().filter(|&g| usable(g, prefix)).collect();
86            match usable.len() {
87                0 => return always(),
88                1 => arms.push(usable[0]),
89                _ => arms.push(bool_expr(pg_sys::BoolExprType::AND_EXPR, usable)),
90            }
91        }
92        match arms.len() {
93            0 => always(),
94            1 => arms[0],
95            _ => bool_expr(pg_sys::BoolExprType::OR_EXPR, arms),
96        }
97    }
98}
99
100struct Found {
101    prefix: i16,
102    /// Per call column: each occurrence's conditions, innermost last.
103    occurrences: Vec<Vec<Vec<*mut pg_sys::Node>>>,
104}
105
106impl Found {
107    unsafe fn visit(&mut self, node: *mut pg_sys::Node, guards: &[*mut pg_sys::Node]) {
108        use pg_sys::NodeTag::*;
109        unsafe {
110            if node.is_null() {
111                return;
112            }
113            match (*node).type_ {
114                T_Var => {
115                    let var = node.cast::<pg_sys::Var>();
116                    if (*var).varno == pg_sys::INDEX_VAR && (*var).varattno > self.prefix {
117                        self.occurrences[((*var).varattno - self.prefix - 1) as usize].push(guards.to_vec());
118                    }
119                }
120                T_BoolExpr => {
121                    let expr = node.cast::<pg_sys::BoolExpr>();
122                    let before = match (*expr).boolop {
123                        pg_sys::BoolExprType::AND_EXPR => Some(pg_sys::BoolTestType::IS_NOT_FALSE),
124                        pg_sys::BoolExprType::OR_EXPR => Some(pg_sys::BoolTestType::IS_NOT_TRUE),
125                        _ => None,
126                    };
127                    let mut inner = guards.to_vec();
128                    for arm in ptrs::<pg_sys::Node>((*expr).args) {
129                        self.visit(arm, &inner);
130                        if let Some(kind) = before {
131                            inner.push(test(arm, kind));
132                        }
133                    }
134                }
135                T_CaseExpr => {
136                    let case = node.cast::<pg_sys::CaseExpr>();
137                    self.visit((*case).arg.cast(), guards);
138                    let mut inner = guards.to_vec();
139                    for when in ptrs::<pg_sys::CaseWhen>((*case).args) {
140                        let condition = (*when).expr.cast::<pg_sys::Node>();
141                        self.visit(condition, &inner);
142                        let mut chosen = inner.clone();
143                        chosen.push(test(condition, pg_sys::BoolTestType::IS_TRUE));
144                        self.visit((*when).result.cast(), &chosen);
145                        inner.push(test(condition, pg_sys::BoolTestType::IS_NOT_TRUE));
146                    }
147                    self.visit((*case).defresult.cast(), &inner);
148                }
149                // Every other node evaluates all of its children.
150                _ => {
151                    let mut children = Vec::new();
152                    let mut first = true;
153                    walk(node, &mut |n| {
154                        if first {
155                            first = false;
156                            return true;
157                        }
158                        children.push(n);
159                        false
160                    });
161                    for child in children {
162                        self.visit(child, guards);
163                    }
164                }
165            }
166        }
167    }
168}
169
170/// Whether `guard` can be evaluated on the scan tuple before the calls
171/// are answered, with the result Postgres will get.
172unsafe fn usable(guard: *mut pg_sys::Node, prefix: i16) -> bool {
173    use pg_sys::NodeTag::*;
174    let mut ok = true;
175    unsafe {
176        walk(guard, &mut |n| {
177            match (*n).type_ {
178                T_Var => {
179                    let var = n.cast::<pg_sys::Var>();
180                    ok &= (*var).varno == pg_sys::INDEX_VAR && (*var).varattno <= prefix;
181                }
182                T_SubPlan | T_AlternativeSubPlan | T_CaseTestExpr | T_Aggref | T_WindowFunc | T_GroupingFunc
183                | T_PlaceHolderVar => ok = false,
184                _ => {}
185            }
186            ok
187        });
188        ok && !pg_sys::contain_volatile_functions(guard)
189    }
190}
191
192/// `arg IS [NOT] TRUE|FALSE`, over a copy of `arg`.
193unsafe fn test(arg: *mut pg_sys::Node, kind: pg_sys::BoolTestType::Type) -> *mut pg_sys::Node {
194    unsafe {
195        let node = pg_sys::palloc0(size_of::<pg_sys::BooleanTest>()).cast::<pg_sys::BooleanTest>();
196        (*node).xpr.type_ = pg_sys::NodeTag::T_BooleanTest;
197        (*node).arg = pg_sys::copyObjectImpl(arg.cast()).cast();
198        (*node).booltesttype = kind;
199        (*node).location = -1;
200        node.cast()
201    }
202}
203
204unsafe fn bool_expr(op: pg_sys::BoolExprType::Type, args: Vec<*mut pg_sys::Node>) -> *mut pg_sys::Node {
205    unsafe { pg_sys::makeBoolExpr(op, append(std::ptr::null_mut(), args), -1).cast() }
206}
207
208/// The reach of a call with no condition.
209pub(super) unsafe fn always() -> *mut pg_sys::Node {
210    unsafe { pg_sys::makeBoolConst(true, false) }
211}
212
213/// Whether `reach` is [`always`].
214pub(super) unsafe fn is_always(reach: *mut pg_sys::Node) -> bool {
215    unsafe {
216        is_a(reach.cast(), pg_sys::NodeTag::T_Const) && {
217            let c = reach.cast::<pg_sys::Const>();
218            !(*c).constisnull && (*c).constvalue.value() != 0
219        }
220    }
221}