Skip to main content

query/dist_plan/
analyzer.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::{BTreeMap, BTreeSet, HashSet};
16use std::sync::Arc;
17
18use common_telemetry::debug;
19use datafusion::config::{ConfigExtension, ExtensionOptions};
20use datafusion::datasource::DefaultTableSource;
21use datafusion::error::Result as DfResult;
22use datafusion_common::config::ConfigOptions;
23use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRewriter};
24use datafusion_common::{Column, ScalarValue};
25use datafusion_expr::expr::{Exists, InSubquery};
26use datafusion_expr::utils::expr_to_columns;
27use datafusion_expr::{Expr, LogicalPlan, LogicalPlanBuilder, Subquery, col as col_fn};
28use datafusion_optimizer::analyzer::AnalyzerRule;
29use datafusion_optimizer::decorrelate_lateral_join::DecorrelateLateralJoin;
30use datafusion_optimizer::decorrelate_predicate_subquery::DecorrelatePredicateSubquery;
31use datafusion_optimizer::eliminate_filter::EliminateFilter;
32use datafusion_optimizer::extract_equijoin_predicate::ExtractEquijoinPredicate;
33use datafusion_optimizer::filter_null_join_keys::FilterNullJoinKeys;
34use datafusion_optimizer::optimizer::Optimizer;
35use datafusion_optimizer::propagate_empty_relation::PropagateEmptyRelation;
36use datafusion_optimizer::push_down_filter::PushDownFilter;
37use datafusion_optimizer::rewrite_set_comparison::RewriteSetComparison;
38use datafusion_optimizer::scalar_subquery_to_join::ScalarSubqueryToJoin;
39use datafusion_optimizer::simplify_expressions::SimplifyExpressions;
40use promql::extension_plan::SeriesDivide;
41use substrait::{DFLogicalSubstraitConvertor, SubstraitPlan};
42use table::metadata::TableType;
43use table::table::adapter::DfTableProviderAdapter;
44
45use crate::dist_plan::RemoteDynFilterProducerId;
46use crate::dist_plan::analyzer::utils::{
47    PatchOptimizerContext, PlanTreeExpressionSimplifier, aliased_columns_for,
48    rewrite_merge_sort_exprs,
49};
50use crate::dist_plan::commutativity::{
51    Categorizer, Commutativity, partial_commutative_transformer,
52};
53use crate::dist_plan::merge_scan::MergeScanLogicalPlan;
54use crate::dist_plan::merge_sort::MergeSortLogicalPlan;
55use crate::metrics::PUSH_DOWN_FALLBACK_ERRORS_TOTAL;
56use crate::options::ScheduledTimeExtension;
57use crate::plan::ExtractExpr;
58use crate::query_engine::DefaultSerializer;
59
60#[cfg(test)]
61mod test;
62
63mod fallback;
64pub(crate) mod utils;
65
66pub(crate) use utils::AliasMapping;
67
68/// Placeholder for other physical partition columns that are not in logical table
69const OTHER_PHY_PART_COL_PLACEHOLDER: &str = "__OTHER_PHYSICAL_PART_COLS_PLACEHOLDER__";
70
71#[derive(Debug, Clone)]
72pub struct DistPlannerOptions {
73    pub allow_query_fallback: bool,
74}
75
76impl ConfigExtension for DistPlannerOptions {
77    const PREFIX: &'static str = "dist_planner";
78}
79
80impl ExtensionOptions for DistPlannerOptions {
81    fn as_any(&self) -> &dyn std::any::Any {
82        self
83    }
84
85    fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
86        self
87    }
88
89    fn cloned(&self) -> Box<dyn ExtensionOptions> {
90        Box::new(self.clone())
91    }
92
93    fn set(&mut self, key: &str, value: &str) -> DfResult<()> {
94        Err(datafusion_common::DataFusionError::NotImplemented(format!(
95            "DistPlannerOptions does not support set key: {key} with value: {value}"
96        )))
97    }
98
99    fn entries(&self) -> Vec<datafusion::config::ConfigEntry> {
100        vec![datafusion::config::ConfigEntry {
101            key: "allow_query_fallback".to_string(),
102            value: Some(self.allow_query_fallback.to_string()),
103            description: "Allow query fallback to fallback plan rewriter",
104        }]
105    }
106}
107
108#[derive(Debug)]
109pub struct DistPlannerAnalyzer;
110
111impl AnalyzerRule for DistPlannerAnalyzer {
112    fn name(&self) -> &str {
113        "DistPlannerAnalyzer"
114    }
115
116    fn analyze(
117        &self,
118        plan: LogicalPlan,
119        config: &ConfigOptions,
120    ) -> datafusion_common::Result<LogicalPlan> {
121        let mut config = config.clone();
122        // Aligned with the behavior in `datafusion_optimizer::OptimizerContext::new()`.
123        config.optimizer.filter_null_join_keys = true;
124        let config = Arc::new(config);
125        let opt = config.extensions.get::<DistPlannerOptions>();
126        let allow_fallback = opt.map(|o| o.allow_query_fallback).unwrap_or(false);
127
128        // When the query is running under a scheduled Flow context, carry the
129        // logical "now" so that `SimplifyExpressions` does not constant-fold
130        // `now()` into wall-clock literals on the remote sub-plans.
131        let scheduled_time = config
132            .extensions
133            .get::<ScheduledTimeExtension>()
134            .and_then(|ext| ext.scheduled_time);
135
136        let optimizer_context = PatchOptimizerContext {
137            inner: datafusion_optimizer::OptimizerContext::new(),
138            config: config.clone(),
139            scheduled_time,
140        };
141
142        let plan = plan
143            .rewrite_with_subqueries(&mut PlanTreeExpressionSimplifier::new(optimizer_context))?
144            .data;
145        let fallback_plan = plan.clone();
146
147        // Run a filter-focused optimizer subset before MergeScan wraps remote
148        // inputs. MergeScan intentionally hides its remote_input from later
149        // optimizer passes; this pass only normalizes/decorrelates enough for
150        // DataFusion's PushDownFilter to put side-local predicates into scans.
151        // Keep this narrow: rules like PushDownLimit, OptimizeProjections, and
152        // DISTINCT rewrites can change global distributed-planning boundaries.
153        let optimizer_context = PatchOptimizerContext {
154            inner: datafusion_optimizer::OptimizerContext::new(),
155            config: config.clone(),
156            scheduled_time,
157        };
158        let plan = match pre_merge_scan_optimizer().optimize(plan, &optimizer_context, |_, _| {}) {
159            Ok(plan) => plan,
160            Err(err) => {
161                if allow_fallback {
162                    common_telemetry::warn!(err; "Failed to pre-optimize plan, using fallback plan rewriter for plan: {fallback_plan}");
163                    PUSH_DOWN_FALLBACK_ERRORS_TOTAL.inc();
164                    return self.use_fallback(fallback_plan);
165                } else {
166                    return Err(err);
167                }
168            }
169        };
170        let plan = plan
171            .transform_down_with_subqueries(&unwrap_dictionary_literals)?
172            .data;
173
174        let result = match self.try_push_down(plan.clone()) {
175            Ok(plan) => plan,
176            Err(err) => {
177                if allow_fallback {
178                    common_telemetry::warn!(err; "Failed to push down plan, using fallback plan rewriter for plan: {plan}");
179                    // if push down failed, use fallback plan rewriter
180                    PUSH_DOWN_FALLBACK_ERRORS_TOTAL.inc();
181                    self.use_fallback(fallback_plan)?
182                } else {
183                    return Err(err);
184                }
185            }
186        };
187
188        Ok(result)
189    }
190}
191
192/// Builds the small optimizer pre-pass that runs before `MergeScan` wrapping.
193///
194/// This is intentionally not DataFusion's full optimizer. After
195/// `PlanRewriter` wraps remote table scans in `MergeScan`,
196/// `MergeScanLogicalPlan::inputs()` hides `remote_input`, so ordinary optimizer
197/// rules can no longer see into the remote side. The main rule we need here is
198/// `PushDownFilter`: it moves side-local join/filter predicates into
199/// `TableScan.filters`, where region pruning and scan-level pruning can use
200/// them.
201///
202/// The rules before `PushDownFilter` are only the minimum cleanup/rewrite steps
203/// needed to make that filter pushdown safe around subqueries and set
204/// comparisons. For example, `RewriteSetComparison` handles ANY/ALL before they
205/// can become scan filters, and the decorrelation/subquery rules expose
206/// supported predicates as joins/filters instead of leaving raw subquery
207/// expressions under a scan.
208///
209/// Keep this list narrow. Do not add broad plan-shaping rules such as
210/// `PushDownLimit`, projection optimization, DISTINCT rewrites, or join-type
211/// rewrites here: those can change the local/remote distributed boundary or
212/// degrade unrelated planning diagnostics. Such rules belong either before this
213/// analyzer or after distributed planning, not in this pre-MergeScan,
214/// filter-focused pass.
215fn pre_merge_scan_optimizer() -> Optimizer {
216    Optimizer::with_rules(vec![
217        Arc::new(RewriteSetComparison::new()),
218        Arc::new(DecorrelatePredicateSubquery::new()),
219        Arc::new(ScalarSubqueryToJoin::new()),
220        Arc::new(DecorrelateLateralJoin::new()),
221        Arc::new(ExtractEquijoinPredicate::new()),
222        Arc::new(EliminateFilter::new()),
223        Arc::new(PropagateEmptyRelation::new()),
224        Arc::new(FilterNullJoinKeys::default()),
225        Arc::new(PushDownFilter::new()),
226        Arc::new(SimplifyExpressions::new()),
227    ])
228}
229
230fn unwrap_dictionary_literals(plan: LogicalPlan) -> DfResult<Transformed<LogicalPlan>> {
231    // A Values plan derives its schema from its expressions. Rewriting only the expressions would
232    // leave that schema inconsistent, which affects DML planning.
233    if matches!(&plan, LogicalPlan::Values(_)) {
234        return Ok(Transformed::no(plan));
235    }
236
237    plan.map_expressions(|expr| {
238        expr.transform_up(|expr| match expr {
239            Expr::Literal(ScalarValue::Dictionary(_, value), metadata) => {
240                Ok(Transformed::yes(Expr::Literal(*value, metadata)))
241            }
242            _ => Ok(Transformed::no(expr)),
243        })
244    })
245}
246
247impl DistPlannerAnalyzer {
248    /// Try push down as many nodes as possible
249    fn try_push_down(&self, plan: LogicalPlan) -> DfResult<LogicalPlan> {
250        // Use the subquery-aware transform so expression subqueries (including
251        // nested scalar subqueries left in place by DataFusion 55's
252        // `enable_physical_uncorrelated_scalar_subquery`) are visited at every
253        // depth and their table scans get wrapped in `MergeScan`.
254        let plan = plan.transform_with_subqueries(&Self::inspect_plan_with_subquery)?;
255        let mut rewriter = PlanRewriter::new(&plan.data);
256        let result = plan.data.rewrite(&mut rewriter)?.data;
257        Self::assign_merge_scan_remote_dyn_filter_producer_ids(result)
258    }
259
260    /// Use fallback plan rewriter to rewrite the plan and only push down table scan nodes
261    fn use_fallback(&self, plan: LogicalPlan) -> DfResult<LogicalPlan> {
262        let mut rewriter = fallback::FallbackPlanRewriter;
263        let result = plan.rewrite(&mut rewriter)?.data;
264        Self::assign_merge_scan_remote_dyn_filter_producer_ids(result)
265    }
266
267    fn inspect_plan_with_subquery(plan: LogicalPlan) -> DfResult<Transformed<LogicalPlan>> {
268        // Workaround for https://github.com/GreptimeTeam/greptimedb/issues/5469 and https://github.com/GreptimeTeam/greptimedb/issues/5799
269        // FIXME(yingwen): Remove the `Limit` plan once we update DataFusion.
270        if let LogicalPlan::Limit(_) | LogicalPlan::Distinct(_) = &plan {
271            return Ok(Transformed::no(plan));
272        }
273
274        let exprs = plan
275            .expressions_consider_join()
276            .into_iter()
277            .map(|e| e.transform(&Self::transform_subquery).map(|x| x.data))
278            .collect::<DfResult<Vec<_>>>()?;
279
280        // Some plans that are special treated (should not call `with_new_exprs` on them)
281        if !matches!(plan, LogicalPlan::Unnest(_)) {
282            let inputs = plan.inputs().into_iter().cloned().collect::<Vec<_>>();
283            Ok(Transformed::yes(plan.with_new_exprs(exprs, inputs)?))
284        } else {
285            Ok(Transformed::no(plan))
286        }
287    }
288
289    fn transform_subquery(expr: Expr) -> DfResult<Transformed<Expr>> {
290        match expr {
291            Expr::Exists(exists) => Ok(Transformed::yes(Expr::Exists(Exists {
292                subquery: Self::handle_subquery(exists.subquery)?,
293                negated: exists.negated,
294            }))),
295            Expr::InSubquery(in_subquery) => Ok(Transformed::yes(Expr::InSubquery(InSubquery {
296                expr: in_subquery.expr,
297                subquery: Self::handle_subquery(in_subquery.subquery)?,
298                negated: in_subquery.negated,
299            }))),
300            Expr::ScalarSubquery(scalar_subquery) => Ok(Transformed::yes(Expr::ScalarSubquery(
301                Self::handle_subquery(scalar_subquery)?,
302            ))),
303
304            _ => Ok(Transformed::no(expr)),
305        }
306    }
307
308    fn handle_subquery(subquery: Subquery) -> DfResult<Subquery> {
309        let mut rewriter = PlanRewriter::new(&subquery.subquery);
310        let mut rewrote_subquery = subquery
311            .subquery
312            .as_ref()
313            .clone()
314            .rewrite(&mut rewriter)?
315            .data;
316        // Workaround. DF doesn't support the first plan in subquery to be an Extension
317        if matches!(rewrote_subquery, LogicalPlan::Extension(_)) {
318            let output_schema = rewrote_subquery.schema().clone();
319            let project_exprs = output_schema
320                .fields()
321                .iter()
322                .map(|f| col_fn(f.name()))
323                .collect::<Vec<_>>();
324            rewrote_subquery = LogicalPlanBuilder::from(rewrote_subquery)
325                .project(project_exprs)?
326                .build()?;
327        }
328
329        Ok(Subquery {
330            subquery: Arc::new(rewrote_subquery),
331            outer_ref_columns: subquery.outer_ref_columns,
332            spans: Default::default(),
333        })
334    }
335
336    fn assign_merge_scan_remote_dyn_filter_producer_ids(
337        plan: LogicalPlan,
338    ) -> DfResult<LogicalPlan> {
339        let mut assigner = MergeScanRemoteDynFilterProducerIdAssigner::default();
340        Ok(plan.rewrite_with_subqueries(&mut assigner)?.data)
341    }
342}
343
344#[derive(Debug, Default)]
345struct RemoteDynFilterProducerIdAllocator {
346    next_remote_dyn_filter_producer_id: u64,
347}
348
349impl RemoteDynFilterProducerIdAllocator {
350    fn allocate(&mut self) -> RemoteDynFilterProducerId {
351        self.next_remote_dyn_filter_producer_id += 1;
352        RemoteDynFilterProducerId::new(self.next_remote_dyn_filter_producer_id)
353    }
354}
355
356/// Assigns query-local RDF producer ids to visible `MergeScan` nodes after plan rewriting.
357#[derive(Debug, Default)]
358struct MergeScanRemoteDynFilterProducerIdAssigner {
359    remote_dyn_filter_producer_id_allocator: RemoteDynFilterProducerIdAllocator,
360}
361
362impl TreeNodeRewriter for MergeScanRemoteDynFilterProducerIdAssigner {
363    type Node = LogicalPlan;
364
365    fn f_up(&mut self, node: Self::Node) -> DfResult<Transformed<Self::Node>> {
366        let LogicalPlan::Extension(extension) = &node else {
367            return Ok(Transformed::no(node));
368        };
369        let Some(merge_scan) = extension
370            .node
371            .as_any()
372            .downcast_ref::<MergeScanLogicalPlan>()
373        else {
374            return Ok(Transformed::no(node));
375        };
376
377        Ok(Transformed::yes(
378            merge_scan
379                .clone()
380                .with_remote_dyn_filter_producer_id(
381                    self.remote_dyn_filter_producer_id_allocator.allocate(),
382                )
383                .into_logical_plan(),
384        ))
385    }
386}
387
388/// Status of the rewriter to mark if the current pass is expanded
389#[derive(Debug, Default, PartialEq, Eq, PartialOrd, Ord)]
390enum RewriterStatus {
391    #[default]
392    Unexpanded,
393    Expanded,
394}
395
396#[derive(Debug, Default)]
397struct PlanRewriter {
398    /// Whether the whole plan this rewriter walks can be encoded to Substrait.
399    /// Used by [`PlanRewriter::should_expand`] to skip the per-node encoding check.
400    whole_plan_encodable: bool,
401    /// Current level in the tree
402    level: usize,
403    /// Simulated stack for the `rewrite` recursion
404    stack: Vec<(LogicalPlan, usize)>,
405    /// Stages to be expanded, will be added as parent node of merge scan one by one
406    stage: Vec<LogicalPlan>,
407    status: RewriterStatus,
408    /// Partition columns of the table in current pass
409    partition_cols: Option<AliasMapping>,
410    /// use stack count as scope to determine column requirements is needed or not
411    /// i.e for a logical plan like:
412    /// ```ignore
413    /// 1: Projection: t.number
414    /// 2: Sort: t.pk1+t.pk2
415    /// 3. Projection: t.number, t.pk1, t.pk2
416    /// ```
417    /// `Sort` will make a column requirement for `t.pk1+t.pk2` at level 2.
418    /// Which making `Projection` at level 1 need to add a ref to `t.pk1` as well.
419    /// So that the expanded plan will be
420    /// ```ignore
421    /// Projection: t.number
422    ///   MergeSort: t.pk1+t.pk2
423    ///     MergeScan: remote_input=
424    /// Projection: t.number, "t.pk1+t.pk2" <--- the original `Projection` at level 1 get added with `t.pk1+t.pk2`
425    ///  Sort: t.pk1+t.pk2
426    ///    Projection: t.number, t.pk1, t.pk2
427    /// ```
428    /// Making `MergeSort` can have `t.pk1+t.pk2` as input.
429    /// Meanwhile `Projection` at level 3 doesn't need to add any new column because 3 > 2
430    /// and col requirements at level 2 is not applicable for level 3.
431    ///
432    /// see more details in test `expand_proj_step_aggr` and `expand_proj_sort_proj`
433    ///
434    /// TODO(discord9): a simpler solution to track column requirements for merge scan
435    column_requirements: Vec<(HashSet<Column>, usize)>,
436    /// Whether to expand on next call
437    /// This is used to handle the case where a plan is transformed, but need to be expanded from it's
438    /// parent node. For example a Aggregate plan is split into two parts in frontend and datanode, and need
439    /// to be expanded from the parent node of the Aggregate plan.
440    expand_on_next_call: bool,
441    /// Expanding on next partial/conditional/transformed commutative plan
442    /// This is used to handle the case where a plan is transformed, but still
443    /// need to push down as many node as possible before next partial/conditional/transformed commutative
444    /// plan. I.e.
445    /// ```ignore
446    /// Limit:
447    ///     Sort:
448    /// ```
449    /// where `Limit` is partial commutative, and `Sort` is conditional commutative.
450    /// In this case, we need to expand the `Limit` plan,
451    /// so that we can push down the `Sort` plan as much as possible.
452    expand_on_next_part_cond_trans_commutative: bool,
453    new_child_plan: Option<LogicalPlan>,
454}
455
456impl PlanRewriter {
457    /// `plan` must be the root of the tree that is about to be rewritten, otherwise
458    /// [`PlanRewriter::should_expand`] may skip a check it has to perform.
459    fn new(plan: &LogicalPlan) -> Self {
460        Self {
461            whole_plan_encodable: DFLogicalSubstraitConvertor
462                .encode(plan, DefaultSerializer)
463                .is_ok(),
464            ..Default::default()
465        }
466    }
467
468    fn get_parent(&self) -> Option<&LogicalPlan> {
469        // level starts from 1, it's safe to minus by 1
470        self.stack
471            .iter()
472            .rev()
473            .find(|(_, level)| *level == self.level - 1)
474            .map(|(node, _)| node)
475    }
476
477    /// Return true if should stop and expand. The input plan is the parent node of current node
478    fn should_expand(&mut self, plan: &LogicalPlan) -> DfResult<bool> {
479        debug!(
480            "Check should_expand at level: {}  with Stack:\n{}, ",
481            self.level,
482            self.stack
483                .iter()
484                .map(|(p, l)| format!("{l}:{}{}", "  ".repeat(l - 1), p.display()))
485                .collect::<Vec<String>>()
486                .join("\n"),
487        );
488        // Substrait encoding recurses from the root to the leaves, so a root that encodes
489        // proves that every node below it encodes as well. `plan` here always comes from
490        // `self.stack`, which holds untouched sub-trees of that root, so one check on the
491        // root covers the whole descent. Each check builds a fresh `SessionState`, which
492        // is the dominant cost on deep PromQL plans.
493        // When the root does not encode, the plan holds at least one node that has to stay
494        // on the frontend and only the per-node check locates it.
495        if !self.whole_plan_encodable
496            && let Err(e) = DFLogicalSubstraitConvertor.encode(plan, DefaultSerializer)
497        {
498            debug!(
499                "PlanRewriter: plan cannot be converted to substrait with error={e:?}, expanding now: {plan}"
500            );
501            return Ok(true);
502        }
503
504        if self.expand_on_next_call {
505            self.expand_on_next_call = false;
506            debug!("PlanRewriter: expand_on_next_call is true, expanding now");
507            return Ok(true);
508        }
509
510        if self.expand_on_next_part_cond_trans_commutative {
511            let comm = Categorizer::check_plan(plan, self.partition_cols.clone())?;
512            match comm {
513                Commutativity::PartialCommutative => {
514                    // a small difference is that for partial commutative, we still need to
515                    // push down it(so `Limit` can be pushed down)
516
517                    // notice how limit needed to be expanded as well to make sure query is correct
518                    // i.e. `Limit fetch=10` need to be pushed down to the leaf node
519                    self.expand_on_next_part_cond_trans_commutative = false;
520                    self.expand_on_next_call = true;
521                }
522                Commutativity::ConditionalCommutative(_)
523                | Commutativity::TransformedCommutative { .. } => {
524                    // again a new node that can be push down, we should just
525                    // do push down now and avoid further expansion
526                    self.expand_on_next_part_cond_trans_commutative = false;
527                    debug!(
528                        "PlanRewriter: meet a new conditional/transformed commutative plan, expanding now: {plan}"
529                    );
530                    return Ok(true);
531                }
532                _ => (),
533            }
534        }
535
536        match Categorizer::check_plan(plan, self.partition_cols.clone())? {
537            Commutativity::Commutative => {
538                // PATCH: we should reconsider SORT's commutativity instead of doing this trick.
539                // explain: for a fully commutative SeriesDivide, its child Sort plan only serves it. I.e., that
540                //   Sort plan is also fully commutative, instead of conditional commutative. So we can remove
541                //   the generated MergeSort from stage safely.
542                if let LogicalPlan::Extension(ext_a) = plan
543                    && ext_a.node.name() == SeriesDivide::name()
544                    && let Some(LogicalPlan::Extension(ext_b)) = self.stage.last()
545                    && ext_b.node.name() == MergeSortLogicalPlan::name()
546                {
547                    // revert last `ConditionalCommutative` result for Sort plan in this case.
548                    // also need to remove any column requirements made by the Sort Plan
549                    // as it may refer to columns later no longer exist(rightfully) like by aggregate or projection
550                    self.stage.pop();
551                    self.expand_on_next_part_cond_trans_commutative = false;
552                    self.column_requirements.clear();
553                }
554            }
555            Commutativity::PartialCommutative => {
556                if let Some(plan) = partial_commutative_transformer(plan) {
557                    // notice this plan is parent of current node, so `self.level - 1` when updating column requirements
558                    self.update_column_requirements(&plan, self.level - 1);
559                    self.expand_on_next_part_cond_trans_commutative = true;
560                    self.stage.push(plan)
561                }
562            }
563            Commutativity::ConditionalCommutative(transformer) => {
564                if let Some(transformer) = transformer
565                    && let Some(plan) = transformer(plan)
566                {
567                    // notice this plan is parent of current node, so `self.level - 1` when updating column requirements
568                    self.update_column_requirements(&plan, self.level - 1);
569                    self.expand_on_next_part_cond_trans_commutative = true;
570                    self.stage.push(plan)
571                }
572            }
573            Commutativity::TransformedCommutative { transformer } => {
574                if let Some(transformer) = transformer {
575                    let transformer_actions = transformer(plan)?;
576                    debug!(
577                        "PlanRewriter: transformed plan: {}\n from {plan}",
578                        transformer_actions
579                            .extra_parent_plans
580                            .iter()
581                            .enumerate()
582                            .map(|(i, p)| format!(
583                                "Extra {i}-th parent plan from parent to child = {}",
584                                p.display()
585                            ))
586                            .collect::<Vec<_>>()
587                            .join("\n")
588                    );
589                    if let Some(new_child_plan) = &transformer_actions.new_child_plan {
590                        debug!("PlanRewriter: new child plan: {}", new_child_plan);
591                    }
592                    if let Some(last_stage) = transformer_actions.extra_parent_plans.last() {
593                        // update the column requirements from the last stage
594                        // notice current plan's parent plan is where we need to apply the column requirements
595                        self.update_column_requirements(last_stage, self.level - 1);
596                    }
597                    self.stage
598                        .extend(transformer_actions.extra_parent_plans.into_iter().rev());
599                    self.expand_on_next_call = true;
600                    self.new_child_plan = transformer_actions.new_child_plan;
601                }
602            }
603            Commutativity::NonCommutative
604            | Commutativity::Unimplemented
605            | Commutativity::Unsupported => {
606                debug!("PlanRewriter: meet a non-commutative plan, expanding now: {plan}");
607                return Ok(true);
608            }
609        }
610
611        Ok(false)
612    }
613
614    /// Update the column requirements for the current plan, plan_level is the level of the plan
615    /// in the stack, which is used to determine if the column requirements are applicable
616    /// for other plans in the stack.
617    fn update_column_requirements(&mut self, plan: &LogicalPlan, plan_level: usize) {
618        debug!(
619            "PlanRewriter: update column requirements for plan: {plan}\n with old column_requirements: {:?}",
620            self.column_requirements
621        );
622        let mut container = HashSet::new();
623        for expr in plan.expressions() {
624            // this method won't fail
625            let _ = expr_to_columns(&expr, &mut container);
626        }
627
628        self.column_requirements.push((container, plan_level));
629        debug!(
630            "PlanRewriter: updated column requirements: {:?}",
631            self.column_requirements
632        );
633    }
634
635    fn is_expanded(&self) -> bool {
636        self.status == RewriterStatus::Expanded
637    }
638
639    fn set_expanded(&mut self) {
640        self.status = RewriterStatus::Expanded;
641    }
642
643    fn set_unexpanded(&mut self) {
644        self.status = RewriterStatus::Unexpanded;
645    }
646
647    fn maybe_set_partitions(&mut self, plan: &LogicalPlan) -> DfResult<()> {
648        if let Some(part_cols) = &mut self.partition_cols {
649            // update partition alias
650            let child = plan.inputs().first().cloned().ok_or_else(|| {
651                datafusion_common::DataFusionError::Internal(format!(
652                    "PlanRewriter: maybe_set_partitions: plan has no child: {plan}"
653                ))
654            })?;
655
656            for (_col_name, alias_set) in part_cols.iter_mut() {
657                let aliased_cols = aliased_columns_for(
658                    &alias_set.clone().into_iter().collect(),
659                    plan,
660                    Some(child),
661                )?;
662                *alias_set = aliased_cols.into_values().flatten().collect();
663            }
664
665            debug!(
666                "PlanRewriter: maybe_set_partitions: updated partition columns: {:?} at plan: {}",
667                part_cols,
668                plan.display()
669            );
670
671            return Ok(());
672        }
673
674        if let LogicalPlan::TableScan(table_scan) = plan
675            && let Some(source) = table_scan.source.downcast_ref::<DefaultTableSource>()
676            && let Some(provider) = source
677                .table_provider
678                .downcast_ref::<DfTableProviderAdapter>()
679        {
680            let table = provider.table();
681            if table.table_type() == TableType::Base {
682                let info = table.table_info();
683                let partition_key_indices = info.meta.partition_key_indices.clone();
684                let schema = info.meta.schema.clone();
685                let mut partition_cols = partition_key_indices
686                    .iter()
687                    .map(|index| schema.column_name_by_index(*index).to_string())
688                    .collect::<Vec<String>>();
689                debug!(
690                    "PlanRewriter: loaded table partition metadata, table: {}, table_id: {}, partition_key_indices: {:?}, partition_columns: {:?}",
691                    info.name, info.ident.table_id, info.meta.partition_key_indices, partition_cols,
692                );
693
694                let partition_rules = table.partition_rules();
695                let exist_phy_part_cols_not_in_logical_table = partition_rules
696                    .map(|r| !r.extra_phy_cols_not_in_logical_table.is_empty())
697                    .unwrap_or(false);
698
699                if exist_phy_part_cols_not_in_logical_table && partition_cols.is_empty() {
700                    // there are other physical partition columns that are not in logical table and part cols are empty
701                    // so we need to add a placeholder for it to prevent certain optimization
702                    // this is used to make sure the final partition columns(that optimizer see) are not empty
703                    // notice if originally partition_cols is not empty, then there is no need to add this place holder,
704                    // as subset of phy part cols can still be used for certain optimization, and it works as if
705                    // those columns are always null
706                    // This helps with distinguishing between non-partitioned table and partitioned table with all phy part cols not in logical table
707                    partition_cols.push(OTHER_PHY_PART_COL_PLACEHOLDER.to_string());
708                }
709                self.partition_cols = Some(
710                            partition_cols
711                                .into_iter()
712                                .map(|c| {
713                                    if c == OTHER_PHY_PART_COL_PLACEHOLDER {
714                                        // for placeholder, just return a empty alias
715                                        return Ok((c.clone(), BTreeSet::new()));
716                                    }
717                                    let index =
718                                        if let Some(c) = plan.schema().index_of_column_by_name(None, &c){
719                                            c
720                                        } else {
721                                            // the `projection` field of `TableScan` doesn't contain the partition columns,
722                                            // this is similar to not having a alias, hence return empty alias set
723                                            return Ok((c.clone(), BTreeSet::new()))
724                                        };
725                                    let column = plan.schema().columns().get(index).cloned().ok_or_else(|| {
726                                        datafusion_common::DataFusionError::Internal(format!(
727                                            "PlanRewriter: maybe_set_partitions: column index {index} out of bounds in schema of plan: {plan}"
728                                        ))
729                                    })?;
730                                    Ok((c.clone(), BTreeSet::from([column])))
731                                })
732                                .collect::<DfResult<AliasMapping>>()?,
733                        );
734            }
735        }
736
737        Ok(())
738    }
739
740    /// pop one stack item and reduce the level by 1
741    fn pop_stack(&mut self) {
742        self.level -= 1;
743        self.stack.pop();
744    }
745
746    fn expand(&mut self, mut on_node: LogicalPlan) -> DfResult<LogicalPlan> {
747        // store schema before expand, new child plan might have a different schema, so not using it
748        let schema = on_node.schema().clone();
749        if let Some(new_child_plan) = self.new_child_plan.take() {
750            // if there is a new child plan, use it as the new root
751            on_node = new_child_plan;
752        }
753        let mut rewriter = EnforceDistRequirementRewriter::new(
754            std::mem::take(&mut self.column_requirements),
755            self.level,
756        );
757        debug!(
758            "PlanRewriter: enforce column requirements for node: {on_node} with rewriter: {rewriter:?}"
759        );
760        on_node = on_node.rewrite(&mut rewriter)?.data;
761        debug!(
762            "PlanRewriter: after enforced column requirements with rewriter: {rewriter:?} for node:\n{on_node}"
763        );
764
765        debug!(
766            "PlanRewriter: expand on node: {on_node} with partition col alias mapping: {:?}",
767            self.partition_cols
768        );
769
770        // add merge scan as the new root
771        let mut node = MergeScanLogicalPlan::new(
772            on_node.clone(),
773            false,
774            // at this stage, the partition cols should be set
775            // treat it as non-partitioned if None
776            self.partition_cols.clone().unwrap_or_default(),
777        )
778        .into_logical_plan();
779
780        // expand stages
781        for new_stage in self.stage.drain(..) {
782            // tracking alias for merge sort's sort exprs
783            let new_stage = if let LogicalPlan::Extension(ext) = &new_stage
784                && let Some(merge_sort) = ext.node.as_any().downcast_ref::<MergeSortLogicalPlan>()
785            {
786                // TODO(discord9): change `on_node` to `node` once alias tracking is supported for merge scan
787                rewrite_merge_sort_exprs(merge_sort, &on_node)?
788            } else {
789                new_stage
790            };
791            node = new_stage
792                .with_new_exprs(new_stage.expressions_consider_join(), vec![node.clone()])?;
793        }
794        self.set_expanded();
795
796        // recover the schema, this make sure after expand the schema is the same as old node
797        // because after expand the raw top node might have extra columns i.e. sorting columns for `Sort` node
798        let node = LogicalPlanBuilder::from(node)
799            .project(schema.iter().map(|(qualifier, field)| {
800                Expr::Column(Column::new(qualifier.cloned(), field.name()))
801            }))?
802            .build()?;
803
804        Ok(node)
805    }
806}
807
808/// Implementation of the [`TreeNodeRewriter`] trait which is responsible for rewriting
809/// logical plans to enforce various requirement for distributed query.
810///
811/// Requirements enforced by this rewriter:
812/// - Enforce column requirements for `LogicalPlan::Projection` nodes. Makes sure the
813///   required columns are available in the sub plan.
814///
815#[derive(Debug)]
816struct EnforceDistRequirementRewriter {
817    /// only enforce column requirements after the expanding node in question,
818    /// meaning only for node with `cur_level` <= `level` will consider adding those column requirements
819    /// TODO(discord9): a simpler solution to track column requirements for merge scan
820    column_requirements: Vec<(HashSet<Column>, usize)>,
821    /// only apply column requirements >= `cur_level`
822    /// this is used to avoid applying column requirements that are not needed
823    /// for the current node, i.e. the node is not in the scope of the column requirements
824    /// i.e, for this plan:
825    /// ```ignore
826    /// Aggregate: min(t.number)
827    ///   Projection: t.number
828    /// ```
829    /// when on `Projection` node, we don't need to apply the column requirements of `Aggregate` node
830    /// because the `Projection` node is not in the scope of the `Aggregate` node
831    cur_level: usize,
832    plan_per_level: BTreeMap<usize, LogicalPlan>,
833}
834
835impl EnforceDistRequirementRewriter {
836    fn new(column_requirements: Vec<(HashSet<Column>, usize)>, cur_level: usize) -> Self {
837        debug!(
838            "Create EnforceDistRequirementRewriter with column_requirements: {:?} at cur_level: {}",
839            column_requirements, cur_level
840        );
841        Self {
842            column_requirements,
843            cur_level,
844            plan_per_level: BTreeMap::new(),
845        }
846    }
847
848    /// Return a mapping from (original column, level) to aliased columns in current node of all
849    /// applicable column requirements
850    /// i.e. only column requirements with level >= `cur_level` will be considered
851    fn get_current_applicable_column_requirements(
852        &self,
853        node: &LogicalPlan,
854    ) -> DfResult<BTreeMap<(Column, usize), BTreeSet<Column>>> {
855        let col_req_per_level = self
856            .column_requirements
857            .iter()
858            .filter(|(_, level)| *level >= self.cur_level)
859            .collect::<Vec<_>>();
860
861        // track alias for columns and use aliased columns instead
862        // aliased col reqs at current level
863        let mut result_alias_mapping = BTreeMap::new();
864        let Some(child) = node.inputs().first().cloned() else {
865            return Ok(Default::default());
866        };
867        for (col_req, level) in col_req_per_level {
868            if let Some(original) = self.plan_per_level.get(level) {
869                // query for alias in current plan
870                let aliased_cols =
871                    aliased_columns_for(&col_req.iter().cloned().collect(), node, Some(original))?;
872                for original_col in col_req {
873                    let aliased_cols = aliased_cols.get(original_col).cloned();
874                    if let Some(cols) = aliased_cols
875                        && !cols.is_empty()
876                    {
877                        result_alias_mapping.insert((original_col.clone(), *level), cols);
878                    } else {
879                        // if no aliased column found in current node, there should be alias in child node as promised by enforce col reqs
880                        // because it should insert required columns in child node
881                        // so we can find the alias in child node
882                        // if not found, it's an internal error
883                        let aliases_in_child = aliased_columns_for(
884                            &[original_col.clone()].into(),
885                            child,
886                            Some(original),
887                        )?;
888                        let Some(aliases) = aliases_in_child
889                            .get(original_col)
890                            .cloned()
891                            .filter(|a| !a.is_empty())
892                        else {
893                            return Err(datafusion_common::DataFusionError::Internal(format!(
894                                "EnforceDistRequirementRewriter: no alias found for required column {original_col} at level {level} in current node's child plan: \n{child} from original plan: \n{original}",
895                            )));
896                        };
897
898                        result_alias_mapping.insert((original_col.clone(), *level), aliases);
899                    }
900                }
901            }
902        }
903        Ok(result_alias_mapping)
904    }
905}
906
907impl TreeNodeRewriter for EnforceDistRequirementRewriter {
908    type Node = LogicalPlan;
909
910    fn f_down(&mut self, node: Self::Node) -> DfResult<Transformed<Self::Node>> {
911        // check that node doesn't have multiple children, i.e. join/subquery
912        if node.inputs().len() > 1 {
913            return Err(datafusion_common::DataFusionError::Internal(
914                "EnforceDistRequirementRewriter: node with multiple inputs is not supported"
915                    .to_string(),
916            ));
917        }
918        self.plan_per_level.insert(self.cur_level, node.clone());
919        self.cur_level += 1;
920        Ok(Transformed::no(node))
921    }
922
923    fn f_up(&mut self, node: Self::Node) -> DfResult<Transformed<Self::Node>> {
924        self.cur_level -= 1;
925        // first get all applicable column requirements
926
927        // make sure all projection applicable scope has the required columns
928        if let LogicalPlan::Projection(ref projection) = node {
929            let mut applicable_column_requirements =
930                self.get_current_applicable_column_requirements(&node)?;
931
932            debug!(
933                "EnforceDistRequirementRewriter: applicable column requirements at level {} = {:?} for node {}",
934                self.cur_level,
935                applicable_column_requirements,
936                node.display()
937            );
938
939            for expr in &projection.expr {
940                let (qualifier, name) = expr.qualified_name();
941                let column = Column::new(qualifier, name);
942                applicable_column_requirements.retain(|_col_level, alias_set| {
943                    // remove all columns that are already in the projection exprs
944                    !alias_set.contains(&column)
945                });
946            }
947            if applicable_column_requirements.is_empty() {
948                return Ok(Transformed::no(node));
949            }
950
951            let mut new_exprs = projection.expr.clone();
952            for (col, alias_set) in &applicable_column_requirements {
953                // use the first alias in alias set as the column to add
954                new_exprs.push(Expr::Column(alias_set.first().cloned().ok_or_else(
955                    || {
956                        datafusion_common::DataFusionError::Internal(
957                            format!("EnforceDistRequirementRewriter: alias set is empty, for column {col:?} in node {node}"),
958                        )
959                    },
960                )?));
961            }
962            let new_node =
963                node.with_new_exprs(new_exprs, node.inputs().into_iter().cloned().collect())?;
964            debug!(
965                "EnforceDistRequirementRewriter: added missing columns {:?} to projection node from old node: \n{node}\n Making new node: \n{new_node}",
966                applicable_column_requirements
967            );
968
969            // update plan for later use
970            self.plan_per_level.insert(self.cur_level, new_node.clone());
971
972            // still need to continue for next projection if applicable
973            return Ok(Transformed::yes(new_node));
974        }
975        Ok(Transformed::no(node))
976    }
977}
978
979impl TreeNodeRewriter for PlanRewriter {
980    type Node = LogicalPlan;
981
982    /// descend
983    fn f_down<'a>(&mut self, node: Self::Node) -> DfResult<Transformed<Self::Node>> {
984        self.level += 1;
985        self.stack.push((node.clone(), self.level));
986        // decendening will clear the stage
987        self.stage.clear();
988        self.set_unexpanded();
989        self.partition_cols = None;
990        Ok(Transformed::no(node))
991    }
992
993    /// ascend
994    ///
995    /// Besure to call `pop_stack` before returning
996    fn f_up(&mut self, node: Self::Node) -> DfResult<Transformed<Self::Node>> {
997        // only expand once on each ascending
998        if self.is_expanded() {
999            self.pop_stack();
1000            return Ok(Transformed::no(node));
1001        }
1002
1003        // only expand when the leaf is table scan
1004        if node.inputs().is_empty() && !matches!(node, LogicalPlan::TableScan(_)) {
1005            self.set_expanded();
1006            self.pop_stack();
1007            return Ok(Transformed::no(node));
1008        }
1009
1010        self.maybe_set_partitions(&node)?;
1011
1012        let Some(parent) = self.get_parent() else {
1013            debug!("Plan Rewriter: expand now for no parent found for node: {node}");
1014            let node = self.expand(node);
1015            debug!(
1016                "PlanRewriter: expanded plan: {}",
1017                match &node {
1018                    Ok(n) => n.to_string(),
1019                    Err(e) => format!("Error expanding plan: {e}"),
1020                }
1021            );
1022            let node = node?;
1023            self.pop_stack();
1024            return Ok(Transformed::yes(node));
1025        };
1026
1027        let parent = parent.clone();
1028
1029        if self.should_expand(&parent)? {
1030            // TODO(ruihang): does this work for nodes with multiple children?;
1031            debug!(
1032                "PlanRewriter: should expand child:\n {node}\n Of Parent: {}",
1033                parent.display()
1034            );
1035            let node = self.expand(node);
1036            debug!(
1037                "PlanRewriter: expanded plan: {}",
1038                match &node {
1039                    Ok(n) => n.to_string(),
1040                    Err(e) => format!("Error expanding plan: {e}"),
1041                }
1042            );
1043            let node = node?;
1044            self.pop_stack();
1045            return Ok(Transformed::yes(node));
1046        }
1047
1048        self.pop_stack();
1049        Ok(Transformed::no(node))
1050    }
1051}