1use std::collections::VecDeque;
4use std::fmt;
5use std::sync::{Arc, Mutex, MutexGuard};
6
7use datafusion::common::config::ConfigOptions;
8use datafusion::common::JoinType;
9use datafusion::common::stats::Precision;
10use datafusion::common::tree_node::{Transformed, TreeNode};
11use datafusion::common::Result;
12use datafusion::physical_optimizer::sanity_checker::SanityCheckPlan;
13use datafusion::physical_optimizer::PhysicalOptimizerRule;
14use datafusion::physical_plan::aggregates::{AggregateExec, AggregateMode};
15use datafusion::physical_plan::coalesce_partitions::CoalescePartitionsExec;
16use datafusion::physical_expr::expressions::Column;
17use datafusion::physical_plan::filter::FilterExec;
18use datafusion::physical_plan::joins::HashJoinExec;
19use datafusion::physical_plan::projection::ProjectionExec;
20use datafusion::physical_plan::repartition::RepartitionExec;
21use datafusion::physical_plan::sorts::sort::SortExec;
22use datafusion::physical_plan::sorts::sort_preserving_merge::SortPreservingMergeExec;
23use datafusion::physical_plan::{
24 displayable, Distribution, ExecutionPlan, ExecutionPlanProperties, Partitioning, StatisticsArgs,
25 StatisticsContext,
26};
27
28use crate::exec::{MetalExec, MetalOp};
29use crate::translate;
30
31#[derive(Debug, Clone)]
40#[non_exhaustive]
41pub struct ArrowMetalConfig {
42 pub min_rows: usize,
44 pub accept_inexact: bool,
46 pub take_when_unknown: bool,
48 pub sort: bool,
50 pub topk: bool,
52 pub aggregate: bool,
54 pub filter: bool,
56 pub aggregate_choice: AggregateChoice,
58 pub table_rows: Option<usize>,
61 pub join: bool,
63 pub join_choice: JoinChoice,
65 pub report_plans: usize,
69}
70
71#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
80pub enum AggregateChoice {
81 #[default]
85 Measured,
86 ArrowMetal,
88 DataFusion,
90}
91
92#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
94#[non_exhaustive]
95pub enum JoinChoice {
96 #[default]
99 Measured,
100 ArrowMetal,
102}
103
104impl Default for ArrowMetalConfig {
137 fn default() -> Self {
138 Self {
139 min_rows: 250_000,
140 accept_inexact: false,
141 take_when_unknown: false,
142 sort: true,
143 topk: false,
144 aggregate: true,
145 filter: false,
146 aggregate_choice: AggregateChoice::Measured,
147 table_rows: None,
148 join: true,
149 join_choice: JoinChoice::Measured,
150 report_plans: 64,
151 }
152 }
153}
154
155impl ArrowMetalConfig {
156 pub fn all() -> Self {
161 Self { topk: true, aggregate: true, filter: true, ..Self::default() }
162 }
163
164 pub fn with_min_rows(mut self, rows: usize) -> Self {
166 self.min_rows = rows;
167 self
168 }
169 pub fn with_accept_inexact(mut self, on: bool) -> Self {
171 self.accept_inexact = on;
172 self
173 }
174 pub fn with_take_when_unknown(mut self, on: bool) -> Self {
176 self.take_when_unknown = on;
177 self
178 }
179 pub fn with_sort(mut self, on: bool) -> Self {
181 self.sort = on;
182 self
183 }
184 pub fn with_topk(mut self, on: bool) -> Self {
186 self.topk = on;
187 self
188 }
189 pub fn with_aggregate(mut self, on: bool) -> Self {
191 self.aggregate = on;
192 self
193 }
194 pub fn with_filter(mut self, on: bool) -> Self {
196 self.filter = on;
197 self
198 }
199 pub fn with_aggregate_choice(mut self, choice: AggregateChoice) -> Self {
201 self.aggregate_choice = choice;
202 self
203 }
204 pub fn with_table_rows(mut self, rows: Option<usize>) -> Self {
206 self.table_rows = rows;
207 self
208 }
209 pub fn with_join(mut self, on: bool) -> Self {
211 self.join = on;
212 self
213 }
214 pub fn with_join_choice(mut self, choice: JoinChoice) -> Self {
216 self.join_choice = choice;
217 self
218 }
219 pub fn with_report_plans(mut self, plans: usize) -> Self {
221 self.report_plans = plans;
222 self
223 }
224}
225
226#[derive(Debug, Clone, PartialEq)]
228#[non_exhaustive]
229pub struct Decision {
230 pub node: String,
232 pub taken: bool,
235 pub reason: String,
237 pub runtime_fallback: bool,
239 pub groups: Option<GroupChoice>,
242}
243
244#[derive(Debug, Clone, PartialEq)]
246#[non_exhaustive]
247pub struct GroupChoice {
248 pub rows: usize,
250 pub estimate: Option<crate::probe::GroupEstimate>,
252}
253
254impl Decision {
255 pub(crate) fn runtime_fallback(op: &MetalOp, msg: &str) -> Self {
256 let reason = if msg.starts_with(crate::gpu::DATA_DEPENDENT) {
257 format!("{msg}; ran the DataFusion plan instead")
258 } else {
259 format!("ArrowMetal error at run time, ran the DataFusion plan instead: {msg}")
260 };
261 Decision { node: format!("MetalExec {op:?}"), taken: false, reason, runtime_fallback: true, groups: None }
262 }
263
264 pub(crate) fn memory_hand_back(op: &MetalOp, what: &str, err: &str) -> Self {
265 Decision {
266 node: format!("MetalExec {op:?}"),
267 taken: false,
268 reason: format!("the memory pool refused the reservation for {what}, ran the DataFusion plan instead: {err}"),
269 runtime_fallback: true,
270 groups: None,
271 }
272 }
273
274 pub(crate) fn runtime_choice(
275 op: &MetalOp,
276 on_arrowmetal: bool,
277 reason: String,
278 rows: usize,
279 estimate: Option<crate::probe::GroupEstimate>,
280 ) -> Self {
281 let reason = format!("{}: {reason}", if on_arrowmetal { "ran on ArrowMetal" } else { "handed back to DataFusion" });
282 Decision {
283 node: format!("MetalExec {op:?}"),
284 taken: on_arrowmetal,
285 reason,
286 runtime_fallback: false,
287 groups: Some(GroupChoice { rows, estimate }),
288 }
289 }
290
291 pub fn is_runtime_choice(&self) -> bool {
293 self.groups.is_some()
294 }
295
296 pub fn is_data_dependent(&self) -> bool {
298 self.runtime_fallback && self.reason.starts_with(crate::gpu::DATA_DEPENDENT)
299 }
300}
301
302impl fmt::Display for Decision {
303 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
304 let tag = if self.runtime_fallback {
305 "FALLBACK"
306 } else if self.groups.is_some() {
307 if self.taken { "GPU" } else { "HANDBACK" }
308 } else if self.taken {
309 "TAKEN"
310 } else {
311 "LEFT"
312 };
313 write!(f, "{tag:8} {} -- {}", self.node, self.reason)
314 }
315}
316
317#[derive(Debug, Clone, Default)]
319#[non_exhaustive]
320pub struct Report(Vec<Decision>);
321
322impl Report {
323 pub fn decisions(&self) -> &[Decision] {
325 &self.0
326 }
327 pub fn taken(&self) -> impl Iterator<Item = &Decision> {
329 self.0.iter().filter(|d| d.taken && d.groups.is_none())
330 }
331 pub fn left(&self) -> impl Iterator<Item = &Decision> {
333 self.0.iter().filter(|d| !d.taken && !d.runtime_fallback && d.groups.is_none())
334 }
335 pub fn runtime_choices(&self) -> impl Iterator<Item = &Decision> {
337 self.0.iter().filter(|d| d.groups.is_some())
338 }
339 pub fn runtime_fallbacks(&self) -> impl Iterator<Item = &Decision> {
342 self.0.iter().filter(|d| d.runtime_fallback)
343 }
344}
345
346impl fmt::Display for Report {
347 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
348 for d in &self.0 {
349 writeln!(f, "{d}")?;
350 }
351 Ok(())
352 }
353}
354
355#[derive(Debug)]
357pub(crate) struct Log {
358 entries: VecDeque<(u64, Decision)>,
359 plan: u64,
360 keep: u64,
361}
362
363impl Log {
364 fn new(keep: usize) -> Self {
365 Self { entries: VecDeque::new(), plan: 0, keep: keep.max(1) as u64 }
366 }
367
368 fn begin_plan(&mut self) -> u64 {
370 self.plan += 1;
371 let first = self.plan.saturating_sub(self.keep - 1);
372 while self.entries.front().is_some_and(|(p, _)| *p < first) {
373 self.entries.pop_front();
374 }
375 self.plan
376 }
377
378 pub(crate) fn push(&mut self, plan: u64, d: Decision) {
380 if plan + self.keep > self.plan {
381 self.entries.push_back((plan, d));
382 }
383 }
384}
385
386pub(crate) type SharedLog = Arc<Mutex<Log>>;
387
388pub(crate) fn lock(log: &SharedLog) -> MutexGuard<'_, Log> {
391 log.lock().unwrap_or_else(|p| p.into_inner())
392}
393
394#[derive(Debug, Clone)]
398pub struct ArrowMetalRule {
399 config: ArrowMetalConfig,
400 log: SharedLog,
401}
402
403struct Candidate {
404 op: std::result::Result<MetalOp, String>,
405 input: Arc<dyn ExecutionPlan>,
407 right: Option<Arc<dyn ExecutionPlan>>,
409 original: Option<Arc<dyn ExecutionPlan>>,
411 reparent: Option<Arc<dyn ExecutionPlan>>,
413}
414
415impl Candidate {
416 fn new(op: std::result::Result<MetalOp, String>, input: Arc<dyn ExecutionPlan>) -> Self {
417 Self { op, input, right: None, original: None, reparent: None }
418 }
419}
420
421fn below_exchanges(p: &Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
425 let mut cur = Arc::clone(p);
426 loop {
427 let next = if let Some(r) = cur.downcast_ref::<RepartitionExec>() {
428 Arc::clone(r.input())
429 } else if let Some(c) = cur.downcast_ref::<CoalescePartitionsExec>() {
430 if c.fetch().is_some() {
431 return cur;
432 }
433 Arc::clone(c.input())
434 } else {
435 return cur;
436 };
437 cur = next;
438 }
439}
440
441fn map_through_projection(
444 spm_expr: &datafusion::physical_expr::LexOrdering,
445 proj: &ProjectionExec,
446) -> Option<Vec<(usize, arrow::compute::SortOptions)>> {
447 let mut out = Vec::new();
448 for s in spm_expr.iter() {
449 let c = s.expr.downcast_ref::<Column>()?;
450 let pe = proj.expr().get(c.index())?;
451 let inner = pe.expr.downcast_ref::<Column>()?;
452 out.push((inner.index(), s.options));
453 }
454 Some(out)
455}
456
457fn sort_keys(expr: &datafusion::physical_expr::LexOrdering) -> Option<Vec<(usize, arrow::compute::SortOptions)>> {
458 expr.iter().map(|s| s.expr.downcast_ref::<Column>().map(|c| (c.index(), s.options))).collect()
459}
460
461impl ArrowMetalRule {
462 pub fn new(config: ArrowMetalConfig) -> Self {
464 let keep = config.report_plans;
465 Self { config, log: Arc::new(Mutex::new(Log::new(keep))) }
466 }
467
468 pub fn config(&self) -> &ArrowMetalConfig {
470 &self.config
471 }
472
473 pub fn report(&self) -> Report {
476 Report(lock(&self.log).entries.iter().map(|(_, d)| d.clone()).collect())
477 }
478
479 pub fn clear_report(&self) {
481 lock(&self.log).entries.clear();
482 }
483
484 fn record(&self, plan: u64, node: &Arc<dyn ExecutionPlan>, taken: bool, reason: String) {
485 let node = displayable(node.as_ref()).one_line().to_string().trim_end().to_string();
486 lock(&self.log).push(plan, Decision { node, taken, reason, runtime_fallback: false, groups: None });
487 }
488
489 fn sort_disabled(&self, fetch: Option<usize>) -> Option<String> {
491 match fetch {
492 None if !self.config.sort => Some("sort disabled in config".into()),
493 Some(n) if !self.config.topk => Some(format!("top-k (sort with fetch {n}) disabled in config")),
494 _ => None,
495 }
496 }
497
498 fn candidate(&self, node: &Arc<dyn ExecutionPlan>) -> Option<Candidate> {
501 if let Some(spm) = node.downcast_ref::<SortPreservingMergeExec>() {
502 let (sort_node, proj) = if spm.input().downcast_ref::<SortExec>().is_some() {
507 (Arc::clone(spm.input()), None)
508 } else {
509 let p = spm.input().downcast_ref::<ProjectionExec>()?;
510 p.input().downcast_ref::<SortExec>()?;
511 (Arc::clone(p.input()), Some(Arc::clone(spm.input())))
512 };
513 let sort = sort_node.downcast_ref::<SortExec>()?;
514 let fetch = match (spm.fetch(), sort.fetch()) {
515 (Some(a), Some(b)) => Some(a.min(b)),
516 (a, b) => a.or(b),
517 };
518 if let Some(why) = self.sort_disabled(fetch) {
519 return Some(Candidate::new(Err(why), Arc::clone(sort.input())));
520 }
521 let same_order = match &proj {
522 None => sort.expr() == spm.expr(),
523 Some(p) => {
524 let p = p.downcast_ref::<ProjectionExec>()?;
525 let mapped = map_through_projection(spm.expr(), p);
526 mapped.is_some() && mapped == sort_keys(sort.expr())
527 }
528 };
529 if !same_order {
530 return Some(Candidate::new(
531 Err("merge ordering differs from the sort's".into()),
532 Arc::clone(sort.input()),
533 ));
534 }
535 let input = Arc::clone(sort.input());
536 let op = translate::sort_op(sort.expr(), &input.schema(), fetch);
537 let mut c = Candidate::new(op, input);
538 if proj.is_some() {
539 c.original = Some(Arc::new(
541 SortPreservingMergeExec::new(sort.expr().clone(), Arc::clone(&sort_node)).with_fetch(fetch),
542 ));
543 c.reparent = proj;
544 }
545 return Some(c);
546 }
547 if let Some(sort) = node.downcast_ref::<SortExec>() {
548 let input = Arc::clone(sort.input());
549 if let Some(why) = self.sort_disabled(sort.fetch()) {
550 return Some(Candidate::new(Err(why), input));
551 }
552 if sort.preserve_partitioning() && input.output_partitioning().partition_count() > 1 {
553 return Some(Candidate::new(Err("per-partition sort (preserve_partitioning) with no replaced merge above it".into()), input));
554 }
555 let op = translate::sort_op(sort.expr(), &input.schema(), sort.fetch());
556 return Some(Candidate::new(op, input));
557 }
558 if let Some(agg) = node.downcast_ref::<AggregateExec>() {
559 if !self.config.aggregate {
560 return Some(Candidate::new(Err("aggregate disabled in config".into()), Arc::clone(agg.input())));
561 }
562 if node.output_ordering().is_some() {
563 return Some(Candidate::new(
564 Err("the aggregate's output carries an ordering (sorted input), which the GPU group-by does not keep".into()),
565 Arc::clone(agg.input()),
566 ));
567 }
568 return Some(match agg.mode() {
569 AggregateMode::Single | AggregateMode::SinglePartitioned => {
570 Candidate::new(translate::aggregate_op(agg), Arc::clone(agg.input()))
571 }
572 AggregateMode::Final | AggregateMode::FinalPartitioned => {
573 let mut cur = Arc::clone(agg.input());
576 loop {
577 if cur.downcast_ref::<RepartitionExec>().is_some()
578 || cur.downcast_ref::<CoalescePartitionsExec>().is_some()
579 {
580 let next = Arc::clone(cur.children()[0]);
581 cur = next;
582 continue;
583 }
584 break;
585 }
586 match cur.downcast_ref::<AggregateExec>() {
587 Some(p) if *p.mode() == AggregateMode::Partial && translate::same_aggregates(agg, p) => {
588 Candidate::new(translate::aggregate_op(p), Arc::clone(p.input()))
589 }
590 _ => Candidate::new(Err("Final aggregate without a matching Partial below its exchange".into()), Arc::clone(agg.input())),
591 }
592 }
593 AggregateMode::Partial => Candidate::new(Err("Partial aggregate whose Final was not replaced".into()), Arc::clone(agg.input())),
594 AggregateMode::PartialReduce => Candidate::new(Err("PartialReduce aggregate".into()), Arc::clone(agg.input())),
595 });
596 }
597 if let Some(f) = node.downcast_ref::<FilterExec>() {
598 let input = Arc::clone(f.input());
599 if !self.config.filter {
600 return Some(Candidate::new(Err("filter disabled in config".into()), input));
601 }
602 if node.fetch().is_some() {
603 return Some(Candidate::new(Err("filter with a fetch limit".into()), input));
604 }
605 if node.output_ordering().is_some() && input.output_partitioning().partition_count() > 1 {
606 return Some(Candidate::new(Err("order-preserving filter over several partitions".into()), input));
607 }
608 let projection = f.projection().as_ref().map(|p| p.iter().copied().collect::<Vec<usize>>());
609 let op = translate::filter_op(f.predicate(), &input.schema(), projection);
610 return Some(Candidate::new(op, input));
611 }
612 if let Some(j) = node.downcast_ref::<HashJoinExec>() {
613 let (left, right) = (below_exchanges(j.left()), below_exchanges(j.right()));
614 let ordered_ok = matches!(j.join_type(), JoinType::Inner | JoinType::Right)
619 && Arc::ptr_eq(&right, j.right())
620 && j.right().output_partitioning().partition_count() == 1;
621 let op = if !self.config.join {
622 Err("join disabled in config".into())
623 } else if node.output_ordering().is_some() && !ordered_ok {
624 Err("the join's output carries an ordering of a probe side split over partitions".into())
625 } else if Arc::ptr_eq(&left, &right) {
626 Err("both join inputs are one plan node".into())
627 } else {
628 translate::join_op(j)
629 };
630 let mut c = Candidate::new(op, left);
631 c.right = Some(right);
632 return Some(c);
633 }
634 None
635 }
636
637 fn statistics_rows(input: &Arc<dyn ExecutionPlan>) -> Precision<usize> {
639 match StatisticsContext::new().compute(input.as_ref(), &StatisticsArgs::new()) {
640 Ok(s) => s.num_rows,
641 Err(_) => Precision::Absent,
642 }
643 }
644
645 fn row_count(&self, input: &Arc<dyn ExecutionPlan>) -> (Option<usize>, String) {
648 match Self::statistics_rows(input) {
649 Precision::Exact(n) => (Some(n), format!("{n} (exact)")),
650 Precision::Inexact(n) if self.config.accept_inexact => (Some(n), format!("~{n} (inexact, accepted)")),
651 Precision::Inexact(n) => (None, format!("~{n} (an estimate; accept_inexact is off)")),
652 Precision::Absent => (None, "unknown".into()),
653 }
654 }
655
656 fn size_ok(&self, input: &Arc<dyn ExecutionPlan>) -> (bool, String, Option<usize>) {
659 let rows = Self::statistics_rows(input);
660 let min = self.config.min_rows;
661 match rows {
662 Precision::Exact(n) => (n >= min, format!("input rows {n} (exact) vs min_rows {min}"), Some(n)),
663 Precision::Inexact(n) if self.config.accept_inexact => {
664 (n >= min, format!("input rows ~{n} (inexact, accepted) vs min_rows {min}"), Some(n))
665 }
666 Precision::Inexact(n) => (false, format!("input rows ~{n} are an estimate (accept_inexact is off)"), None),
667 Precision::Absent => (
668 self.config.take_when_unknown,
669 format!("input row count unknown (take_when_unknown = {})", self.config.take_when_unknown),
670 None,
671 ),
672 }
673 }
674
675 fn visit(&self, plan: u64, node: Arc<dyn ExecutionPlan>, top: usize) -> Result<Transformed<Arc<dyn ExecutionPlan>>> {
678 if node.downcast_ref::<MetalExec>().is_some() {
679 return Ok(Transformed::no(node));
680 }
681 let Some(c) = self.candidate(&node) else {
682 return Ok(Transformed::no(node));
683 };
684 let op = match c.op {
685 Ok(op) => op,
686 Err(why) => {
687 self.record(plan, &node, false, why);
688 return Ok(Transformed::no(node));
689 }
690 };
691 let at_top = Arc::as_ptr(&node) as *const () as usize == top;
696 let wrap = match node.output_partitioning() {
697 p if p.partition_count() <= 1 => None,
698 _ if at_top => None,
699 Partitioning::Hash(exprs, n) => Some(Partitioning::Hash(exprs.clone(), *n)),
700 Partitioning::RoundRobinBatch(n) | Partitioning::UnknownPartitioning(n) => {
701 Some(Partitioning::RoundRobinBatch(*n))
702 }
703 Partitioning::Range(_) => {
704 self.record(plan, &node, false, "range-partitioned output".into());
705 return Ok(Transformed::no(node));
706 }
707 };
708 let mut input = Arc::clone(&c.input);
711 while let Some(r) = input.downcast_ref::<RepartitionExec>() {
712 if !matches!(r.partitioning(), Partitioning::RoundRobinBatch(_)) {
713 break;
714 }
715 let next = Arc::clone(r.input());
716 input = next;
717 }
718 let mut inputs = vec![Arc::clone(&input)];
719 let (ok, mut size, rows) = match &c.right {
720 None => self.size_ok(&input),
721 Some(right) => {
722 inputs.push(Arc::clone(right));
724 let (lrows, ltext) = self.row_count(&input);
725 let (rrows, rtext) = self.row_count(right);
726 let rows_text = format!("left (build) rows {ltext}, right (probe) rows {rtext}");
727 match (lrows, rrows) {
728 (Some(l), Some(r)) => {
729 let min = self.config.min_rows;
730 let n = l.max(r);
731 let mut why = format!("{rows_text}; the larger vs min_rows {min}");
732 let mut ok = n >= min;
733 if ok && self.config.join_choice == JoinChoice::Measured {
734 match crate::choice::join_takes(&op, input.as_ref(), right.as_ref(), l as u64, r as u64) {
735 Ok(w) => why = format!("{why}; {w}"),
736 Err(w) => {
737 why = format!("{why}; {w}");
738 ok = false;
739 }
740 }
741 }
742 (ok, why, Some(l + r))
743 }
744 _ => (false, format!("{rows_text}: a join needs both row counts"), None),
745 }
746 }
747 };
748 if !ok {
749 self.record(plan, &node, false, size);
750 return Ok(Transformed::no(node));
751 }
752 if matches!(op, MetalOp::Aggregate { .. }) && self.config.aggregate_choice == AggregateChoice::Measured {
755 let Some(shape) = crate::choice::shape(&op, input.as_ref()) else {
756 self.record(plan, &node, false, "aggregate without a shape".into());
757 return Ok(Transformed::no(node));
758 };
759 if let Some(n) = self.config.table_rows.or(rows) {
760 match crate::choice::any_bucket(&shape, n as u64) {
761 Ok(why) => size = format!("{size}; {why}"),
762 Err(why) => {
763 self.record(plan, &node, false, format!("{size}; {why}"));
764 return Ok(Transformed::no(node));
765 }
766 }
767 }
768 }
769 let original = c.original.unwrap_or_else(|| Arc::clone(&node));
770 let mut wrap = wrap;
771 let mut out_parts = 1;
772 if c.right.is_some() && !at_top {
773 out_parts = node.output_partitioning().partition_count();
776 if !matches!(wrap, Some(Partitioning::Hash(..))) {
777 wrap = None;
778 }
779 }
780 let metal: Arc<dyn ExecutionPlan> = Arc::new(
781 MetalExec::new(
782 op,
783 inputs,
784 original,
785 (Arc::clone(&self.log), plan),
786 crate::exec::AggSettings {
787 choice: self.config.aggregate_choice,
788 table_rows: self.config.table_rows,
789 rows_hint: rows,
790 },
791 )
792 .with_output_partitions(out_parts),
793 );
794 let mut reason = size;
795 if out_parts > 1 {
796 reason.push_str(&format!("; output in {out_parts} partitions"));
797 }
798 if at_top && node.output_partitioning().partition_count() > 1 {
799 reason.push_str("; output kept at one partition (only projections above it)");
800 }
801 let metal = match c.reparent {
802 #[allow(deprecated)] Some(p) => {
804 reason.push_str("; projection kept above it");
805 p.with_new_children(vec![metal])?
806 }
807 None => metal,
808 };
809 let out: Arc<dyn ExecutionPlan> = match wrap {
810 None => metal,
811 Some(p) => {
812 if c.right.is_some() && matches!(p, Partitioning::Hash(..)) {
813 reason.push_str(&format!(
814 "; output re-partitioned to {p} where the parent requires it (to keep the parent's distribution)"
815 ));
816 } else {
817 reason.push_str(&format!("; output re-partitioned to {p} to keep the parent's distribution"));
818 }
819 Arc::new(RepartitionExec::try_new(metal, p)?)
820 }
821 };
822 self.record(plan, &node, true, reason);
823 Ok(Transformed::yes(out))
824 }
825}
826
827#[allow(deprecated)] fn drop_hash_where_not_required(node: Arc<dyn ExecutionPlan>) -> Result<Transformed<Arc<dyn ExecutionPlan>>> {
833 let required = node.required_input_distribution();
834 let children: Vec<Arc<dyn ExecutionPlan>> = node.children().into_iter().cloned().collect();
835 let mut changed = false;
836 let mut new_children = Vec::with_capacity(children.len());
837 for (i, c) in children.into_iter().enumerate() {
838 let swap = match c.downcast_ref::<RepartitionExec>() {
839 Some(r) => {
840 let over_join = r
841 .input()
842 .downcast_ref::<MetalExec>()
843 .is_some_and(|m| matches!(m.op(), MetalOp::Join { .. }));
844 match r.partitioning() {
845 Partitioning::Hash(..)
846 if over_join && matches!(required.get(i), Some(Distribution::UnspecifiedDistribution)) =>
847 {
848 Some(Arc::clone(r.input()))
849 }
850 _ => None,
851 }
852 }
853 None => None,
854 };
855 match swap {
856 Some(s) => {
857 changed = true;
858 new_children.push(s);
859 }
860 None => new_children.push(c),
861 }
862 }
863 if !changed {
864 return Ok(Transformed::no(node));
865 }
866 Ok(Transformed::yes(node.with_new_children(new_children)?))
867}
868
869impl PhysicalOptimizerRule for ArrowMetalRule {
870 fn optimize(&self, plan: Arc<dyn ExecutionPlan>, config: &ConfigOptions) -> Result<Arc<dyn ExecutionPlan>> {
871 let mut top = Arc::clone(&plan);
872 while let Some(p) = top.downcast_ref::<ProjectionExec>() {
873 let next = Arc::clone(p.input());
874 top = next;
875 }
876 let top = Arc::as_ptr(&top) as *const () as usize;
877 let id = lock(&self.log).begin_plan();
878 let plan = plan.transform_down(|n| self.visit(id, n, top))?.data;
879 let relaxed = Arc::clone(&plan).transform_up(drop_hash_where_not_required)?;
883 if relaxed.transformed && SanityCheckPlan::new().optimize(Arc::clone(&relaxed.data), config).is_ok() {
884 return Ok(relaxed.data);
885 }
886 Ok(plan)
887 }
888
889 fn name(&self) -> &str {
890 "ArrowMetalRule"
891 }
892
893 fn schema_check(&self) -> bool {
894 true
895 }
896}