Skip to main content

query/optimizer/
global_limit.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::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                // The query limit is the only initial limit.
324                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}