1use 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
68const 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 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 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 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 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
192fn 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 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 fn try_push_down(&self, plan: LogicalPlan) -> DfResult<LogicalPlan> {
250 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 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 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 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 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#[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#[derive(Debug, Default, PartialEq, Eq, PartialOrd, Ord)]
390enum RewriterStatus {
391 #[default]
392 Unexpanded,
393 Expanded,
394}
395
396#[derive(Debug, Default)]
397struct PlanRewriter {
398 whole_plan_encodable: bool,
401 level: usize,
403 stack: Vec<(LogicalPlan, usize)>,
405 stage: Vec<LogicalPlan>,
407 status: RewriterStatus,
408 partition_cols: Option<AliasMapping>,
410 column_requirements: Vec<(HashSet<Column>, usize)>,
436 expand_on_next_call: bool,
441 expand_on_next_part_cond_trans_commutative: bool,
453 new_child_plan: Option<LogicalPlan>,
454}
455
456impl PlanRewriter {
457 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 self.stack
471 .iter()
472 .rev()
473 .find(|(_, level)| *level == self.level - 1)
474 .map(|(node, _)| node)
475 }
476
477 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 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 self.expand_on_next_part_cond_trans_commutative = false;
520 self.expand_on_next_call = true;
521 }
522 Commutativity::ConditionalCommutative(_)
523 | Commutativity::TransformedCommutative { .. } => {
524 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 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 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 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 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 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 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 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 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 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 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 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 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 let schema = on_node.schema().clone();
749 if let Some(new_child_plan) = self.new_child_plan.take() {
750 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 let mut node = MergeScanLogicalPlan::new(
772 on_node.clone(),
773 false,
774 self.partition_cols.clone().unwrap_or_default(),
777 )
778 .into_logical_plan();
779
780 for new_stage in self.stage.drain(..) {
782 let new_stage = if let LogicalPlan::Extension(ext) = &new_stage
784 && let Some(merge_sort) = ext.node.as_any().downcast_ref::<MergeSortLogicalPlan>()
785 {
786 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 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#[derive(Debug)]
816struct EnforceDistRequirementRewriter {
817 column_requirements: Vec<(HashSet<Column>, usize)>,
821 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 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 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 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 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 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 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 !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 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 self.plan_per_level.insert(self.cur_level, new_node.clone());
971
972 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 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 self.stage.clear();
988 self.set_unexpanded();
989 self.partition_cols = None;
990 Ok(Transformed::no(node))
991 }
992
993 fn f_up(&mut self, node: Self::Node) -> DfResult<Transformed<Self::Node>> {
997 if self.is_expanded() {
999 self.pop_stack();
1000 return Ok(Transformed::no(node));
1001 }
1002
1003 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 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}