1use std::sync::Arc;
16
17use datafusion::config::ConfigOptions;
18use datafusion::physical_optimizer::PhysicalOptimizerRule;
19use datafusion::physical_plan::coalesce_partitions::CoalescePartitionsExec;
20use datafusion::physical_plan::filter::FilterExec;
21use datafusion::physical_plan::limit::GlobalLimitExec;
22use datafusion::physical_plan::repartition::RepartitionExec;
23use datafusion::physical_plan::sorts::sort_preserving_merge::SortPreservingMergeExec;
24use datafusion::physical_plan::{
25 ChildrenPropertiesMode, ExecutionPlan, ExecutionPlanProperties, ReplaceChildrenOptions,
26};
27use datafusion_common::Result as DfResult;
28use datafusion_physical_expr::{Distribution, OrderingRequirements, Partitioning};
29
30use crate::dist_plan::MergeSortExec;
31
32#[derive(Debug)]
33pub struct EnsureGlobalLimitForFetch;
34
35impl PhysicalOptimizerRule for EnsureGlobalLimitForFetch {
36 fn optimize(
37 &self,
38 plan: Arc<dyn ExecutionPlan>,
39 _config: &ConfigOptions,
40 ) -> DfResult<Arc<dyn ExecutionPlan>> {
41 Self::optimize_plan(plan, ParentContext::default())
42 }
43
44 fn name(&self) -> &str {
45 "EnsureGlobalLimitForFetch"
46 }
47
48 fn schema_check(&self) -> bool {
49 true
50 }
51}
52
53impl EnsureGlobalLimitForFetch {
54 fn optimize_plan(
55 plan: Arc<dyn ExecutionPlan>,
56 parent: ParentContext,
57 ) -> DfResult<Arc<dyn ExecutionPlan>> {
58 let children = plan.children();
59 let plan = if children.is_empty() {
60 plan
61 } else {
62 let required_input_distribution = plan.input_distribution_requirements();
63 let required_input_ordering = plan.required_input_ordering();
64 let maintains_input_order = plan.maintains_input_order();
65 let child_parent = ParentContext {
66 global_fetch: provided_global_fetch(&plan),
67 required_ordering: None,
68 required_distribution: Distribution::UnspecifiedDistribution,
69 partitioning_to_restore: None,
70 preserve_hash_partitioning: false,
71 };
72 let children = children
73 .into_iter()
74 .enumerate()
75 .map(|(idx, child)| {
76 let required_distribution = required_input_distribution
77 .child_distribution(idx)
78 .cloned()
79 .unwrap_or(Distribution::UnspecifiedDistribution);
80 let partitioning_to_restore =
81 partitioning_to_restore_for(child, &required_distribution)
82 .or_else(|| inherited_partitioning_to_restore(&plan, child, &parent));
83 let preserve_hash_partitioning = partitioning_to_restore.is_some();
84 let required_ordering = required_input_ordering
85 .get(idx)
86 .cloned()
87 .unwrap_or(None)
88 .or_else(|| {
89 maintains_input_order
90 .get(idx)
91 .copied()
92 .unwrap_or(false)
93 .then(|| parent.required_ordering.clone())
94 .flatten()
95 });
96 let parent = ParentContext {
97 required_ordering,
98 required_distribution,
99 partitioning_to_restore,
100 preserve_hash_partitioning,
101 ..child_parent.clone()
102 };
103 Self::optimize_plan(Arc::clone(child), parent)
104 })
105 .collect::<DfResult<Vec<_>>>()?;
106 plan.replace_children(
107 children,
108 ReplaceChildrenOptions::new(ChildrenPropertiesMode::Recompute),
109 )?
110 };
111
112 let Some(fetch) = plan.fetch() else {
113 return Ok(plan);
114 };
115
116 if parent
117 .global_fetch
118 .is_some_and(|parent_fetch| parent_fetch <= fetch)
119 || !plan.is::<FilterExec>()
120 || plan.output_partitioning().partition_count() <= 1
121 {
122 return Ok(plan);
123 }
124
125 add_global_fetch(
126 plan,
127 fetch,
128 parent.required_ordering,
129 parent.partitioning_to_restore,
130 )
131 }
132}
133
134#[derive(Clone)]
135struct ParentContext {
136 global_fetch: Option<usize>,
137 required_ordering: Option<OrderingRequirements>,
138 required_distribution: Distribution,
139 partitioning_to_restore: Option<Partitioning>,
140 preserve_hash_partitioning: bool,
141}
142
143impl Default for ParentContext {
144 fn default() -> Self {
145 Self {
146 global_fetch: None,
147 required_ordering: None,
148 required_distribution: Distribution::UnspecifiedDistribution,
149 partitioning_to_restore: None,
150 preserve_hash_partitioning: false,
151 }
152 }
153}
154
155fn provided_global_fetch(plan: &Arc<dyn ExecutionPlan>) -> Option<usize> {
156 let fetch = plan.fetch()?;
157 (plan.is::<GlobalLimitExec>()
158 || plan.is::<CoalescePartitionsExec>()
159 || plan.is::<SortPreservingMergeExec>()
160 || plan.is::<MergeSortExec>())
161 .then_some(fetch)
162}
163
164fn add_global_fetch(
165 plan: Arc<dyn ExecutionPlan>,
166 fetch: usize,
167 required_ordering: Option<OrderingRequirements>,
168 partitioning_to_restore: Option<Partitioning>,
169) -> DfResult<Arc<dyn ExecutionPlan>> {
170 let plan = if required_ordering.is_some()
171 && let Some(ordering) = plan.output_ordering().cloned()
172 {
173 Arc::new(SortPreservingMergeExec::new(ordering, plan).with_fetch(Some(fetch)))
174 as Arc<dyn ExecutionPlan>
175 } else {
176 Arc::new(CoalescePartitionsExec::new(plan).with_fetch(Some(fetch)))
177 as Arc<dyn ExecutionPlan>
178 };
179
180 restore_required_partitioning(plan, partitioning_to_restore)
181}
182
183fn restore_required_partitioning(
184 plan: Arc<dyn ExecutionPlan>,
185 partitioning_to_restore: Option<Partitioning>,
186) -> DfResult<Arc<dyn ExecutionPlan>> {
187 let Some(partitioning) = partitioning_to_restore else {
188 return Ok(plan);
189 };
190
191 if partitioning.partition_count() <= 1 || !matches!(&partitioning, Partitioning::Hash(_, _)) {
192 return Ok(plan);
193 }
194
195 Ok(Arc::new(
196 RepartitionExec::try_new(plan, partitioning)?.with_preserve_order(),
197 ))
198}
199
200fn partitioning_to_restore_for(
201 child: &Arc<dyn ExecutionPlan>,
202 required_distribution: &Distribution,
203) -> Option<Partitioning> {
204 if !matches!(required_distribution, Distribution::KeyPartitioned(_))
205 || child.output_partitioning().partition_count() <= 1
206 {
207 return None;
208 }
209
210 if child
211 .output_partitioning()
212 .satisfaction(required_distribution, child.equivalence_properties(), false)
213 .is_satisfied()
214 {
215 Some(child.output_partitioning().clone())
216 } else {
217 Some(
218 required_distribution
219 .clone()
220 .create_partitioning(child.output_partitioning().partition_count()),
221 )
222 }
223}
224
225fn inherited_partitioning_to_restore(
226 plan: &Arc<dyn ExecutionPlan>,
227 child: &Arc<dyn ExecutionPlan>,
228 parent: &ParentContext,
229) -> Option<Partitioning> {
230 if child.output_partitioning().partition_count() <= 1
231 || !matches!(child.output_partitioning(), Partitioning::Hash(_, _))
232 || !matches!(plan.output_partitioning(), Partitioning::Hash(_, _))
233 || plan.output_partitioning().partition_count()
234 != child.output_partitioning().partition_count()
235 {
236 return None;
237 }
238
239 let satisfies_parent_distribution = matches!(
240 parent.required_distribution,
241 Distribution::KeyPartitioned(_)
242 ) && plan
243 .output_partitioning()
244 .satisfaction(
245 &parent.required_distribution,
246 plan.equivalence_properties(),
247 false,
248 )
249 .is_satisfied();
250
251 (satisfies_parent_distribution || parent.preserve_hash_partitioning)
252 .then(|| child.output_partitioning().clone())
253}
254
255#[cfg(test)]
256mod tests {
257 use datafusion::arrow::array::{Array, Int32Array};
258 use datafusion::arrow::compute::SortOptions;
259 use datafusion::arrow::datatypes::{DataType, Field, Schema};
260 use datafusion::arrow::record_batch::RecordBatch;
261 use datafusion::execution::TaskContext;
262 use datafusion::physical_expr::expressions::{col, lit};
263 use datafusion::physical_optimizer::optimizer::PhysicalOptimizer;
264 use datafusion::physical_plan::aggregates::{
265 AggregateExec, AggregateMode, LimitOptions, PhysicalGroupBy,
266 };
267 use datafusion::physical_plan::filter::FilterExecBuilder;
268 use datafusion::physical_plan::joins::{HashJoinExec, PartitionMode};
269 use datafusion::physical_plan::limit::{GlobalLimitExec, LocalLimitExec};
270 use datafusion::physical_plan::projection::ProjectionExec;
271 use datafusion::physical_plan::repartition::RepartitionExec;
272 use datafusion::physical_plan::sorts::sort::SortExec;
273 use datafusion::physical_plan::test::TestMemoryExec;
274 use datafusion_common::{JoinType, NullEquality};
275 use datafusion_physical_expr::{LexOrdering, Partitioning, PhysicalSortExpr};
276
277 use super::*;
278
279 async fn optimize_and_collect_twice(
280 mut plan: Arc<dyn ExecutionPlan>,
281 config: &ConfigOptions,
282 ) -> Vec<Vec<i32>> {
283 let mut results = Vec::with_capacity(2);
284 for _ in 0..2 {
285 for rule in PhysicalOptimizer::new().rules {
286 plan = rule.optimize(plan, config).unwrap();
287 }
288 let batches = datafusion::physical_plan::collect(
289 Arc::clone(&plan),
290 Arc::new(TaskContext::default()),
291 )
292 .await
293 .unwrap();
294 results.push(
295 batches
296 .iter()
297 .flat_map(|batch| {
298 batch
299 .column(0)
300 .as_any()
301 .downcast_ref::<Int32Array>()
302 .unwrap()
303 .values()
304 .to_vec()
305 })
306 .collect(),
307 );
308 }
309 results
310 }
311
312 #[tokio::test]
313 async fn physical_optimizer_keeps_global_distinct_limit_across_two_passes() {
314 for soft_limit in [false, true] {
315 let mut config = ConfigOptions::new();
316 config.execution.target_partitions = 3;
317 config.optimizer.enable_distinct_aggregation_soft_limit = soft_limit;
318
319 for mode in [
320 AggregateMode::FinalPartitioned,
321 AggregateMode::SinglePartitioned,
322 ] {
323 let input = input_with_all_hash_partitions();
325 let aggregate = match mode {
326 AggregateMode::FinalPartitioned => {
327 agg(agg(input, AggregateMode::Partial), mode)
328 }
329 AggregateMode::SinglePartitioned => agg(hash_repartition(input), mode),
330 _ => unreachable!(),
331 };
332 let plan =
333 Arc::new(GlobalLimitExec::new(aggregate, 0, Some(1))) as Arc<dyn ExecutionPlan>;
334
335 let results = optimize_and_collect_twice(plan, &config).await;
336 assert_eq!(
337 results.iter().map(Vec::len).collect::<Vec<_>>(),
338 vec![1, 1],
339 "soft limit enabled: {soft_limit}, mode: {mode:?}",
340 );
341 }
342 }
343 }
344
345 #[tokio::test]
346 async fn physical_optimizer_keeps_count_over_distinct_limit_across_two_passes() {
347 use datafusion::datasource::MemTable;
348 use datafusion::execution::context::{SessionConfig, SessionContext as DFSessionContext};
349 use datafusion_common::ScalarValue;
350
351 for soft_limit in [false, true] {
352 let mut config = SessionConfig::new().with_target_partitions(3);
353 config
354 .options_mut()
355 .optimizer
356 .enable_distinct_aggregation_soft_limit = soft_limit;
357 let ctx = DFSessionContext::new_with_config(config.clone());
358 let schema = schema();
359 let batch = RecordBatch::try_new(
360 schema.clone(),
361 vec![Arc::new(Int32Array::from((0..100).collect::<Vec<_>>()))],
362 )
363 .unwrap();
364 let table = MemTable::try_new(schema, vec![vec![batch]; 3]).unwrap();
365 ctx.register_table("t", Arc::new(table)).unwrap();
366 let mut plan = ctx
367 .sql("SELECT COUNT(*) FROM (SELECT DISTINCT a FROM t LIMIT 1) AS limited")
368 .await
369 .unwrap()
370 .create_physical_plan()
371 .await
372 .unwrap();
373
374 for pass in 0..=2 {
375 if pass > 0 {
376 for rule in PhysicalOptimizer::new().rules {
377 plan = rule.optimize(plan, config.options()).unwrap();
378 }
379 }
380 let batches = datafusion::physical_plan::collect(Arc::clone(&plan), ctx.task_ctx())
381 .await
382 .unwrap();
383 let values = batches
384 .iter()
385 .flat_map(|batch| {
386 (0..batch.num_rows()).map(|row| {
387 ScalarValue::try_from_array(batch.column(0).as_ref(), row).unwrap()
388 })
389 })
390 .collect::<Vec<_>>();
391 assert_eq!(
392 values,
393 vec![ScalarValue::Int64(Some(1))],
394 "soft limit enabled: {soft_limit}, additional optimizer passes: {pass}",
395 );
396 }
397 }
398 }
399
400 #[tokio::test]
401 async fn physical_optimizer_keeps_local_limit_per_partition_across_two_passes() {
402 let mut config = ConfigOptions::new();
403 config.execution.target_partitions = 3;
404 let input = input_with_all_hash_partitions();
405 let local_limit = Arc::new(LocalLimitExec::new(
406 agg(hash_repartition(input), AggregateMode::SinglePartitioned),
407 1,
408 )) as Arc<dyn ExecutionPlan>;
409
410 let results = optimize_and_collect_twice(local_limit, &config).await;
411 assert_eq!(results.iter().map(Vec::len).collect::<Vec<_>>(), vec![3, 3]);
412 }
413
414 #[tokio::test]
415 async fn physical_optimizer_keeps_ordered_topk_across_two_passes() {
416 let schema = schema();
417 let losing =
418 RecordBatch::try_new(schema.clone(), vec![Arc::new(Int32Array::from(vec![100]))])
419 .unwrap();
420 let winning =
421 RecordBatch::try_new(schema.clone(), vec![Arc::new(Int32Array::from(vec![1]))])
422 .unwrap();
423 let input = Arc::new(
424 TestMemoryExec::try_new(&[vec![losing], vec![winning]], schema.clone(), None).unwrap(),
425 );
426 let ordering = ordering(schema.as_ref(), false);
427 let topk = Arc::new(
428 SortExec::new(
429 ordering,
430 agg_with_limit_options(
431 input,
432 AggregateMode::FinalPartitioned,
433 Some(LimitOptions::new_with_order(1, false)),
434 ),
435 )
436 .with_fetch(Some(1)),
437 ) as Arc<dyn ExecutionPlan>;
438
439 let mut config = ConfigOptions::new();
440 config.execution.target_partitions = 3;
441 assert_eq!(
442 optimize_and_collect_twice(topk, &config).await,
443 vec![vec![1], vec![1]],
444 );
445 }
446
447 #[tokio::test]
448 async fn physical_optimizer_keeps_ordered_offset_limit_across_two_passes() {
449 let mut config = ConfigOptions::new();
450 config.execution.target_partitions = 3;
451 let offset = Arc::new(GlobalLimitExec::new(
452 Arc::new(SortExec::new(
453 ordering(schema().as_ref(), false),
454 unordered_input(),
455 )),
456 3,
457 Some(1),
458 )) as Arc<dyn ExecutionPlan>;
459
460 assert_eq!(
461 optimize_and_collect_twice(offset, &config).await,
462 vec![vec![2], vec![2]],
463 );
464 }
465
466 #[test]
467 fn adds_global_limit_for_multi_partition_filter_fetch() {
468 let filter = filter_fetch(unordered_input(), 1);
469
470 let optimized =
471 EnsureGlobalLimitForFetch::optimize_plan(filter, ParentContext::default()).unwrap();
472
473 assert!(optimized.is::<CoalescePartitionsExec>());
474 assert_eq!(optimized.fetch(), Some(1));
475 assert_eq!(optimized.output_partitioning().partition_count(), 1);
476 }
477
478 #[test]
479 fn still_visits_subtree_under_global_limit() {
480 let filter = filter_fetch(unordered_input(), 5);
481 let projection = Arc::new(
482 ProjectionExec::try_new(
483 vec![(col("a", filter.schema().as_ref()).unwrap(), "a".to_string())],
484 filter,
485 )
486 .unwrap(),
487 );
488 let limit =
489 Arc::new(GlobalLimitExec::new(projection, 0, Some(10))) as Arc<dyn ExecutionPlan>;
490
491 let optimized =
492 EnsureGlobalLimitForFetch::optimize_plan(limit, ParentContext::default()).unwrap();
493 let projection = optimized.children()[0];
494 let coalesce = projection.children()[0];
495
496 assert!(coalesce.is::<CoalescePartitionsExec>());
497 assert_eq!(coalesce.fetch(), Some(5));
498 }
499
500 #[test]
501 fn keeps_filter_under_parent_global_fetch() {
502 let (input, ordering) = ordered_input();
503 let filter = filter_fetch(input, 1);
504 let merge = Arc::new(SortPreservingMergeExec::new(ordering, filter).with_fetch(Some(1)))
505 as Arc<dyn ExecutionPlan>;
506
507 let optimized =
508 EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
509 let child = optimized.children()[0];
510
511 assert!(optimized.is::<SortPreservingMergeExec>());
512 assert!(child.is::<FilterExec>());
513 }
514
515 #[test]
516 fn adds_tighter_global_fetch_under_looser_parent_fetch() {
517 let (input, ordering) = ordered_input();
518 let filter = filter_fetch(input, 5);
519 let merge = Arc::new(SortPreservingMergeExec::new(ordering, filter).with_fetch(Some(10)))
520 as Arc<dyn ExecutionPlan>;
521
522 let optimized =
523 EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
524 let child = optimized.children()[0];
525
526 assert!(optimized.is::<SortPreservingMergeExec>());
527 assert!(child.is::<SortPreservingMergeExec>());
528 assert_eq!(child.fetch(), Some(5));
529 }
530
531 #[test]
532 fn keeps_filter_under_parent_merge_sort_fetch() {
533 let (input, ordering) = ordered_input();
534 let filter = filter_fetch(input, 1);
535 let merge = merge_sort_fetch(ordering, filter, 1);
536
537 let optimized =
538 EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
539 let child = optimized.children()[0];
540
541 assert!(optimized.is::<MergeSortExec>());
542 assert!(child.is::<FilterExec>());
543 }
544
545 #[test]
546 fn adds_tighter_global_fetch_under_looser_merge_sort_fetch() {
547 let (input, ordering) = ordered_input();
548 let filter = filter_fetch(input, 5);
549 let merge = merge_sort_fetch(ordering, filter, 10);
550
551 let optimized =
552 EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
553 let child = optimized.children()[0];
554
555 assert!(optimized.is::<MergeSortExec>());
556 assert!(child.is::<SortPreservingMergeExec>());
557 assert_eq!(child.fetch(), Some(5));
558 assert!(child.children()[0].is::<FilterExec>());
559 }
560
561 #[test]
562 fn preserves_parent_ordering_requirement() {
563 let (input, ordering) = ordered_input();
564 let filter = filter_fetch(input, 1);
565 let merge =
566 Arc::new(SortPreservingMergeExec::new(ordering, filter)) as Arc<dyn ExecutionPlan>;
567
568 let optimized =
569 EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
570 let child = optimized.children()[0];
571
572 assert!(optimized.is::<SortPreservingMergeExec>());
573 assert!(child.is::<SortPreservingMergeExec>());
574 assert_eq!(child.fetch(), Some(1));
575 }
576
577 #[test]
578 fn uses_child_output_ordering_for_merge() {
579 let schema = schema();
580 let required_ordering = ordering(schema.as_ref(), false);
581 let actual_ordering = ordering(schema.as_ref(), true);
582 let batch = batch(schema.clone());
583 let partitions = vec![vec![batch.clone()], vec![batch.clone()], vec![batch]];
584 let input = TestMemoryExec::try_new(&partitions, schema, None)
585 .unwrap()
586 .try_with_sort_information(vec![actual_ordering.clone()])
587 .unwrap();
588 let filter = filter_fetch(Arc::new(input), 1);
589
590 let optimized = add_global_fetch(
591 filter,
592 1,
593 Some(OrderingRequirements::from(required_ordering)),
594 None,
595 )
596 .unwrap();
597 let merge = optimized.downcast_ref::<SortPreservingMergeExec>().unwrap();
598
599 assert_eq!(merge.expr(), &actual_ordering);
600 }
601
602 #[test]
603 fn preserves_inherited_ordering_requirement_through_projection() {
604 let (input, ordering) = ordered_input();
605 let filter = filter_fetch(input, 1);
606 let projection = Arc::new(
607 ProjectionExec::try_new(
608 vec![(col("a", filter.schema().as_ref()).unwrap(), "a".to_string())],
609 filter,
610 )
611 .unwrap(),
612 );
613 let merge =
614 Arc::new(SortPreservingMergeExec::new(ordering, projection)) as Arc<dyn ExecutionPlan>;
615
616 let optimized =
617 EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
618 let projection = optimized.children()[0];
619 let child = projection.children()[0];
620
621 assert!(optimized.is::<SortPreservingMergeExec>());
622 assert!(projection.is::<ProjectionExec>());
623 assert!(child.is::<SortPreservingMergeExec>());
624 assert_eq!(child.fetch(), Some(1));
625 }
626
627 #[test]
628 fn restores_parent_hash_distribution_after_global_fetch() {
629 let left = filter_fetch(hash_repartition(unordered_input()), 1);
630 let right = hash_repartition(unordered_input());
631 let on = vec![(
632 col("a", left.schema().as_ref()).unwrap(),
633 col("a", right.schema().as_ref()).unwrap(),
634 )];
635 let join = Arc::new(
636 HashJoinExec::try_new(
637 left,
638 right,
639 on,
640 None,
641 &JoinType::Inner,
642 None,
643 PartitionMode::Partitioned,
644 NullEquality::NullEqualsNothing,
645 false,
646 )
647 .unwrap(),
648 ) as Arc<dyn ExecutionPlan>;
649
650 let optimized =
651 EnsureGlobalLimitForFetch::optimize_plan(join, ParentContext::default()).unwrap();
652 let left = optimized.children()[0];
653 let repartition = left.downcast_ref::<RepartitionExec>().unwrap();
654
655 assert!(matches!(
656 repartition.partitioning(),
657 Partitioning::Hash(_, 3)
658 ));
659 assert!(repartition.input().is::<CoalescePartitionsExec>());
660 assert_eq!(repartition.input().fetch(), Some(1));
661 }
662
663 #[test]
664 fn restores_inherited_hash_distribution_through_projection() {
665 let filter = filter_fetch(hash_repartition(unordered_input()), 1);
666 let projection = Arc::new(
667 ProjectionExec::try_new(
668 vec![(col("a", filter.schema().as_ref()).unwrap(), "a".to_string())],
669 filter,
670 )
671 .unwrap(),
672 ) as Arc<dyn ExecutionPlan>;
673 let right = hash_repartition(unordered_input());
674 let on = vec![(
675 col("a", projection.schema().as_ref()).unwrap(),
676 col("a", right.schema().as_ref()).unwrap(),
677 )];
678 let join = Arc::new(
679 HashJoinExec::try_new(
680 projection,
681 right,
682 on,
683 None,
684 &JoinType::Inner,
685 None,
686 PartitionMode::Partitioned,
687 NullEquality::NullEqualsNothing,
688 false,
689 )
690 .unwrap(),
691 ) as Arc<dyn ExecutionPlan>;
692
693 let optimized =
694 EnsureGlobalLimitForFetch::optimize_plan(join, ParentContext::default()).unwrap();
695 let projection = optimized.children()[0];
696 let repartition = projection.children()[0]
697 .downcast_ref::<RepartitionExec>()
698 .unwrap();
699
700 assert!(projection.is::<ProjectionExec>());
701 assert!(matches!(
702 repartition.partitioning(),
703 Partitioning::Hash(_, 3)
704 ));
705 assert!(repartition.input().is::<CoalescePartitionsExec>());
706 assert_eq!(repartition.input().fetch(), Some(1));
707 }
708
709 #[test]
710 fn restores_inherited_hash_distribution_through_multiple_projections() {
711 let filter = filter_fetch(hash_repartition(unordered_input()), 1);
712 let projection = project_a(filter);
713 let projection = project_a(projection);
714 let right = hash_repartition(unordered_input());
715 let on = vec![(
716 col("a", projection.schema().as_ref()).unwrap(),
717 col("a", right.schema().as_ref()).unwrap(),
718 )];
719 let join = Arc::new(
720 HashJoinExec::try_new(
721 projection,
722 right,
723 on,
724 None,
725 &JoinType::Inner,
726 None,
727 PartitionMode::Partitioned,
728 NullEquality::NullEqualsNothing,
729 false,
730 )
731 .unwrap(),
732 ) as Arc<dyn ExecutionPlan>;
733
734 let optimized =
735 EnsureGlobalLimitForFetch::optimize_plan(join, ParentContext::default()).unwrap();
736 let outer_projection = optimized.children()[0];
737 let inner_projection = outer_projection.children()[0];
738 let repartition = inner_projection.children()[0]
739 .downcast_ref::<RepartitionExec>()
740 .unwrap();
741
742 assert!(outer_projection.is::<ProjectionExec>());
743 assert!(inner_projection.is::<ProjectionExec>());
744 assert!(matches!(
745 repartition.partitioning(),
746 Partitioning::Hash(_, 3)
747 ));
748 assert!(repartition.input().is::<CoalescePartitionsExec>());
749 assert_eq!(repartition.input().fetch(), Some(1));
750 }
751
752 fn unordered_input() -> Arc<dyn ExecutionPlan> {
753 let schema = schema();
754 let batch = batch(schema.clone());
755 let partitions = vec![vec![batch.clone()], vec![batch.clone()], vec![batch]];
756 Arc::new(TestMemoryExec::try_new(&partitions, schema, None).unwrap())
757 }
758
759 fn input_with_all_hash_partitions() -> Arc<dyn ExecutionPlan> {
760 let schema = schema();
761 let batch = RecordBatch::try_new(
762 schema.clone(),
763 vec![Arc::new(Int32Array::from((0..100).collect::<Vec<_>>()))],
764 )
765 .unwrap();
766 let partitions = vec![vec![batch.clone()], vec![batch.clone()], vec![batch]];
767 Arc::new(TestMemoryExec::try_new(&partitions, schema, None).unwrap())
768 }
769
770 fn agg(input: Arc<dyn ExecutionPlan>, mode: AggregateMode) -> Arc<dyn ExecutionPlan> {
771 agg_with_limit_options(input, mode, None)
772 }
773
774 fn agg_with_limit_options(
775 input: Arc<dyn ExecutionPlan>,
776 mode: AggregateMode,
777 limit_options: Option<LimitOptions>,
778 ) -> Arc<dyn ExecutionPlan> {
779 let schema = input.schema();
780 let group_by = PhysicalGroupBy::new_single(vec![(
781 col("a", schema.as_ref()).unwrap(),
782 "a".to_string(),
783 )]);
784 Arc::new(
785 AggregateExec::try_new(mode, group_by, vec![], vec![], input, schema)
786 .unwrap()
787 .with_limit_options(limit_options),
788 )
789 }
790
791 fn ordered_input() -> (Arc<dyn ExecutionPlan>, LexOrdering) {
792 let schema = schema();
793 let ordering = ordering(schema.as_ref(), false);
794 let batch = batch(schema.clone());
795 let partitions = vec![vec![batch.clone()], vec![batch.clone()], vec![batch]];
796 let input = TestMemoryExec::try_new(&partitions, schema, None)
797 .unwrap()
798 .try_with_sort_information(vec![ordering.clone()])
799 .unwrap();
800
801 (Arc::new(input), ordering)
802 }
803
804 fn filter_fetch(input: Arc<dyn ExecutionPlan>, fetch: usize) -> Arc<dyn ExecutionPlan> {
805 Arc::new(
806 FilterExecBuilder::new(lit(true), input)
807 .with_fetch(Some(fetch))
808 .build()
809 .unwrap(),
810 )
811 }
812
813 fn merge_sort_fetch(
814 ordering: LexOrdering,
815 input: Arc<dyn ExecutionPlan>,
816 fetch: usize,
817 ) -> Arc<dyn ExecutionPlan> {
818 Arc::new(MergeSortExec::new(ordering, input, Some(fetch)))
819 }
820
821 fn hash_repartition(input: Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
822 let partitioning = Partitioning::Hash(vec![col("a", input.schema().as_ref()).unwrap()], 3);
823 Arc::new(RepartitionExec::try_new(input, partitioning).unwrap())
824 }
825
826 fn project_a(input: Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
827 Arc::new(
828 ProjectionExec::try_new(
829 vec![(col("a", input.schema().as_ref()).unwrap(), "a".to_string())],
830 input,
831 )
832 .unwrap(),
833 )
834 }
835
836 fn schema() -> Arc<Schema> {
837 Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)]))
838 }
839
840 fn batch(schema: Arc<Schema>) -> RecordBatch {
841 RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 2, 3]))]).unwrap()
842 }
843
844 fn ordering(schema: &Schema, descending: bool) -> LexOrdering {
845 LexOrdering::new([PhysicalSortExpr::new(
846 col("a", schema).unwrap(),
847 SortOptions {
848 descending,
849 nulls_first: descending,
850 },
851 )])
852 .unwrap()
853 }
854}