lift.rsannotatedlift.rssource279 lines · 11.4 KB · raw
1//! The plan post-pass: calls the relation scan could not claim, because
2//! an upper node evaluates them (an aggregate's arguments, HAVING, a
3//! window or grouping node's own expressions), get a jev scan of their
4//! own between that node and its child.
5//!
6//! It runs in `planner_hook`, after `standard_planner`, so setrefs has
7//! run: the calls read the child's output as `OUTER_VAR` columns. The new
8//! scan's scan tuple is the child's columns, in order, then one column per
9//! call, exactly as `plan_custom_path` lays it out; so the parent keeps
10//! its references to the child's columns, and each call becomes an
11//! `OUTER_VAR` reference to its column. A call reading anything else (a
12//! Param, a SubPlan, an aggregate) stays where it is, and `guard.rs`
13//! refuses the plan before anything is sent.
14
15use std::ffi::c_char;
16use std::sync::OnceLock;
17
18use pgrx::{pg_guard, pg_sys};
19
20use super::Judge;
21use super::ffi::{append, is_a, len, mutate, ptrs, walk};
22use super::plan::{function_oids, tables};
23
24static PREVIOUS: OnceLock<pg_sys::planner_hook_type> = OnceLock::new();
25
26/// # Safety
27/// Called once, from `_PG_init`.
28pub(super) unsafe fn install<J: Judge>() {
29    unsafe {
30        let _ = PREVIOUS.set(pg_sys::planner_hook);
31        pg_sys::planner_hook = Some(planner::<J>);
32    }
33}
34
35#[pg_guard]
36unsafe extern "C-unwind" fn planner<J: Judge>(
37    parse: *mut pg_sys::Query,
38    query_string: *const c_char,
39    cursor_options: i32,
40    bound_params: pg_sys::ParamListInfo,
41) -> *mut pg_sys::PlannedStmt {
42    // SAFETY: the planner's arguments; the statement it returns is ours
43    // to edit until it is returned.
44    unsafe {
45        let stmt = match PREVIOUS.get() {
46            Some(Some(previous)) => previous(parse, query_string, cursor_options, bound_params),
47            _ => pg_sys::standard_planner(parse, query_string, cursor_options, bound_params),
48        };
49        let Some(oids) = function_oids::<J>() else { return stmt };
50        let mut next_id = max_node_id((*stmt).planTree);
51        for sub in ptrs::<pg_sys::Plan>((*stmt).subplans) {
52            next_id = next_id.max(max_node_id(sub));
53        }
54        let mut lift = Lift { oids: &oids, next_id: next_id + 1 };
55        lift.plan((*stmt).planTree);
56        for sub in ptrs::<pg_sys::Plan>((*stmt).subplans) {
57            lift.plan(sub);
58        }
59        // Every jev scan is now in place, with its final quals and
60        // projection: work out where each call is reached.
61        super::reach::plan((*stmt).planTree, children);
62        for sub in ptrs::<pg_sys::Plan>((*stmt).subplans) {
63            super::reach::plan(sub, children);
64        }
65        stmt
66    }
67}
68
69struct Lift<'a> {
70    oids: &'a [pg_sys::Oid],
71    next_id: i32,
72}
73
74/// The children of `plan` the post-pass descends into. Not a Gather's:
75/// what runs in workers never evaluates a call (they are PARALLEL
76/// RESTRICTED), and a scan must not be put there.
77unsafe fn children(plan: *mut pg_sys::Plan) -> Vec<*mut pg_sys::Plan> {
78    use pg_sys::NodeTag::*;
79    unsafe {
80        let mut found = Vec::new();
81        match (*plan).type_ {
82            T_Gather | T_GatherMerge => return found,
83            T_Append => found.extend(ptrs((*plan.cast::<pg_sys::Append>()).appendplans)),
84            T_MergeAppend => found.extend(ptrs((*plan.cast::<pg_sys::MergeAppend>()).mergeplans)),
85            T_SubqueryScan => found.push((*plan.cast::<pg_sys::SubqueryScan>()).subplan),
86            // A jev scan's child is in custom_plans (its lefttree, when it
87            // has one, is the same node, set only for EXPLAIN).
88            T_CustomScan => return ptrs((*plan.cast::<pg_sys::CustomScan>()).custom_plans).collect(),
89            _ => {}
90        }
91        found.extend([(*plan).lefttree, (*plan).righttree].into_iter().filter(|p| !p.is_null()));
92        found
93    }
94}
95
96unsafe fn max_node_id(plan: *mut pg_sys::Plan) -> i32 {
97    unsafe {
98        if plan.is_null() {
99            return 0;
100        }
101        let mut id = (*plan).plan_node_id;
102        for child in children(plan) {
103            id = id.max(max_node_id(child));
104        }
105        id
106    }
107}
108
109impl Lift<'_> {
110    unsafe fn plan(&mut self, plan: *mut pg_sys::Plan) {
111        use pg_sys::NodeTag::*;
112        unsafe {
113            if plan.is_null() {
114                return;
115            }
116            for child in children(plan) {
117                self.plan(child);
118            }
119            let lifts = match (*plan).type_ {
120                // Grouping sets chain further Agg nodes over one input.
121                T_Agg => (*plan.cast::<pg_sys::Agg>()).chain.is_null(),
122                T_WindowAgg | T_Group | T_Result => true,
123                _ => false,
124            };
125            if !lifts || (*plan).lefttree.is_null() || !(*plan).righttree.is_null() {
126                return;
127            }
128            let mut calls: Vec<*mut pg_sys::FuncExpr> = Vec::new();
129            for list in [(*plan).targetlist, (*plan).qual] {
130                walk(list.cast(), &mut |n| {
131                    if is_a(n.cast(), pg_sys::NodeTag::T_FuncExpr) {
132                        let call = n.cast::<pg_sys::FuncExpr>();
133                        if self.oids.contains(&(*call).funcid) {
134                            if reads_only_outer(n)
135                                && !pg_sys::expression_returns_set(n)
136                                && !calls.iter().any(|&c| pg_sys::equal(c.cast(), call.cast()))
137                            {
138                                calls.push(call);
139                            }
140                            return false;
141                        }
142                    }
143                    true
144                });
145            }
146            if calls.is_empty() {
147                return;
148            }
149            let child = (*plan).lefttree;
150            let prefix = len((*child).targetlist) as i32;
151            let scan = self.scan(child, &calls);
152            (*plan).lefttree = scan.cast();
153            let columns: Vec<(*mut pg_sys::FuncExpr, *mut pg_sys::Var)> = calls
154                .iter()
155                .enumerate()
156                .map(|(k, &call)| (call, column(pg_sys::OUTER_VAR, prefix + k as i32 + 1, call.cast())))
157                .collect();
158            let mut replace = |n: *mut pg_sys::Node| {
159                columns
160                    .iter()
161                    .find(|&&(call, _)| is_a(n.cast(), pg_sys::NodeTag::T_FuncExpr) && pg_sys::equal(n.cast(), call.cast()))
162                    .map(|&(_, var)| pg_sys::copyObjectImpl(var.cast()).cast::<pg_sys::Node>())
163            };
164            (*plan).targetlist = mutate((*plan).targetlist.cast(), &mut replace).cast();
165            (*plan).qual = mutate((*plan).qual.cast(), &mut replace).cast();
166        }
167    }
168
169    /// The jev scan over `child`, computing `calls`, laid out as
170    /// `plan_custom_path` lays it out, in post-setrefs form.
171    unsafe fn scan(&mut self, child: *mut pg_sys::Plan, calls: &[*mut pg_sys::FuncExpr]) -> *mut pg_sys::CustomScan {
172        unsafe {
173            let mut scan_tlist = std::ptr::null_mut();
174            let mut tlist = std::ptr::null_mut();
175            let mut resno: i16 = 0;
176            let mut add = |expr: *mut pg_sys::Expr| {
177                resno += 1;
178                scan_tlist = append(scan_tlist, [pg_sys::makeTargetEntry(expr, resno, std::ptr::null_mut::<c_char>(), false)]);
179                let var = column(pg_sys::INDEX_VAR, i32::from(resno), expr.cast());
180                tlist = append(tlist, [pg_sys::makeTargetEntry(var.cast(), resno, std::ptr::null_mut::<c_char>(), false)]);
181            };
182            // The child's columns, read as OUTER_VAR of the child: EXPLAIN
183            // resolves them through the lefttree below.
184            for (i, entry) in ptrs::<pg_sys::TargetEntry>((*child).targetlist).enumerate() {
185                add(column(pg_sys::OUTER_VAR, i as i32 + 1, (*entry).expr.cast()).cast());
186            }
187            let prefix = len((*child).targetlist) as i32;
188            let mut private = pg_sys::lappend_int(std::ptr::null_mut(), prefix);
189            let mut exprs = std::ptr::null_mut();
190            for &call in calls {
191                add(pg_sys::copyObjectImpl(call.cast()).cast());
192                let index = self.oids.iter().position(|&o| o == (*call).funcid).expect("a claimed call");
193                private = pg_sys::lappend_int(private, index as i32);
194                // Evaluated against the scan tuple, whose first columns are
195                // the child's, in order.
196                for arg in ptrs::<pg_sys::Node>((*call).args) {
197                    let arg = mutate(arg, &mut |n| {
198                        if is_a(n.cast(), pg_sys::NodeTag::T_Var) {
199                            let var = pg_sys::copyObjectImpl(n.cast()).cast::<pg_sys::Var>();
200                            (*var).varno = pg_sys::INDEX_VAR;
201                            return Some(var.cast());
202                        }
203                        None
204                    });
205                    exprs = append(exprs, [arg]);
206                }
207            }
208            exprs = append(exprs, calls.iter().map(|_| super::reach::always()));
209
210            let scan = pg_sys::palloc0(size_of::<pg_sys::CustomScan>()).cast::<pg_sys::CustomScan>();
211            let p = &mut (*scan).scan.plan;
212            p.type_ = pg_sys::NodeTag::T_CustomScan;
213            p.startup_cost = (*child).startup_cost;
214            p.total_cost = (*child).total_cost;
215            p.plan_rows = (*child).plan_rows;
216            p.plan_width = (*child).plan_width;
217            #[cfg(not(feature = "pg17"))]
218            {
219                p.disabled_nodes = (*child).disabled_nodes;
220            }
221            p.parallel_aware = false;
222            p.parallel_safe = false;
223            p.async_capable = false;
224            p.plan_node_id = self.next_id;
225            self.next_id += 1;
226            p.targetlist = tlist;
227            p.qual = std::ptr::null_mut();
228            // Only for EXPLAIN, whose deparser resolves OUTER_VAR through
229            // it; the executor starts the child from custom_plans.
230            p.lefttree = child;
231            p.extParam = pg_sys::bms_copy((*child).extParam);
232            p.allParam = pg_sys::bms_copy((*child).allParam);
233            (*scan).scan.scanrelid = 0;
234            (*scan).flags = pg_sys::CUSTOMPATH_SUPPORT_PROJECTION;
235            (*scan).custom_plans = append(std::ptr::null_mut(), [child]);
236            (*scan).custom_exprs = exprs;
237            (*scan).custom_private = private;
238            (*scan).custom_scan_tlist = scan_tlist;
239            (*scan).methods = &tables().scan;
240            scan
241        }
242    }
243}
244
245/// A Var reading column `attno` of `varno` with `expr`'s type.
246unsafe fn column(varno: i32, attno: i32, expr: *mut pg_sys::Node) -> *mut pg_sys::Var {
247    unsafe {
248        pg_sys::makeVar(
249            varno,
250            attno as i16,
251            pg_sys::exprType(expr),
252            pg_sys::exprTypmod(expr),
253            pg_sys::exprCollation(expr),
254            0,
255        )
256    }
257}
258
259/// Whether `call` reads nothing but its input row: every Var is the
260/// child's output, and there is no Param, SubPlan or aggregate.
261unsafe fn reads_only_outer(call: *mut pg_sys::Node) -> bool {
262    use pg_sys::NodeTag::*;
263    let mut ok = true;
264    unsafe {
265        walk(call, &mut |n| {
266            match (*n).type_ {
267                T_Var => {
268                    let var = n.cast::<pg_sys::Var>();
269                    ok &= (*var).varno == pg_sys::OUTER_VAR && (*var).varlevelsup == 0;
270                }
271                T_Param | T_SubPlan | T_AlternativeSubPlan | T_Aggref | T_WindowFunc | T_GroupingFunc
272                | T_PlaceHolderVar => ok = false,
273                _ => {}
274            }
275            ok
276        });
277    }
278    ok
279}