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::{ExecutionPlan, ExecutionPlanProperties};
25use datafusion_common::Result as DfResult;
26use datafusion_physical_expr::{Distribution, OrderingRequirements, Partitioning};
27
28use crate::dist_plan::MergeSortExec;
29
30#[derive(Debug)]
31pub struct EnsureGlobalLimitForFetch;
32
33impl PhysicalOptimizerRule for EnsureGlobalLimitForFetch {
34    fn optimize(
35        &self,
36        plan: Arc<dyn ExecutionPlan>,
37        _config: &ConfigOptions,
38    ) -> DfResult<Arc<dyn ExecutionPlan>> {
39        Self::optimize_plan(plan, ParentContext::default())
40    }
41
42    fn name(&self) -> &str {
43        "EnsureGlobalLimitForFetch"
44    }
45
46    fn schema_check(&self) -> bool {
47        true
48    }
49}
50
51impl EnsureGlobalLimitForFetch {
52    fn optimize_plan(
53        plan: Arc<dyn ExecutionPlan>,
54        parent: ParentContext,
55    ) -> DfResult<Arc<dyn ExecutionPlan>> {
56        let children = plan.children();
57        let plan = if children.is_empty() {
58            plan
59        } else {
60            let required_input_distribution = plan.required_input_distribution();
61            let required_input_ordering = plan.required_input_ordering();
62            let maintains_input_order = plan.maintains_input_order();
63            let child_parent = ParentContext {
64                global_fetch: provided_global_fetch(&plan),
65                required_ordering: None,
66                required_distribution: Distribution::UnspecifiedDistribution,
67                partitioning_to_restore: None,
68                preserve_hash_partitioning: false,
69            };
70            let children = children
71                .into_iter()
72                .enumerate()
73                .map(|(idx, child)| {
74                    let required_distribution = required_input_distribution
75                        .get(idx)
76                        .cloned()
77                        .unwrap_or(Distribution::UnspecifiedDistribution);
78                    let partitioning_to_restore =
79                        partitioning_to_restore_for(child, &required_distribution)
80                            .or_else(|| inherited_partitioning_to_restore(&plan, child, &parent));
81                    let preserve_hash_partitioning = partitioning_to_restore.is_some();
82                    let required_ordering = required_input_ordering
83                        .get(idx)
84                        .cloned()
85                        .unwrap_or(None)
86                        .or_else(|| {
87                            maintains_input_order
88                                .get(idx)
89                                .copied()
90                                .unwrap_or(false)
91                                .then(|| parent.required_ordering.clone())
92                                .flatten()
93                        });
94                    let parent = ParentContext {
95                        required_ordering,
96                        required_distribution,
97                        partitioning_to_restore,
98                        preserve_hash_partitioning,
99                        ..child_parent.clone()
100                    };
101                    Self::optimize_plan(Arc::clone(child), parent)
102                })
103                .collect::<DfResult<Vec<_>>>()?;
104            plan.with_new_children(children)?
105        };
106
107        let Some(fetch) = plan.fetch() else {
108            return Ok(plan);
109        };
110
111        if parent
112            .global_fetch
113            .is_some_and(|parent_fetch| parent_fetch <= fetch)
114            || !plan.as_any().is::<FilterExec>()
115            || plan.output_partitioning().partition_count() <= 1
116        {
117            return Ok(plan);
118        }
119
120        add_global_fetch(
121            plan,
122            fetch,
123            parent.required_ordering,
124            parent.partitioning_to_restore,
125        )
126    }
127}
128
129#[derive(Clone)]
130struct ParentContext {
131    global_fetch: Option<usize>,
132    required_ordering: Option<OrderingRequirements>,
133    required_distribution: Distribution,
134    partitioning_to_restore: Option<Partitioning>,
135    preserve_hash_partitioning: bool,
136}
137
138impl Default for ParentContext {
139    fn default() -> Self {
140        Self {
141            global_fetch: None,
142            required_ordering: None,
143            required_distribution: Distribution::UnspecifiedDistribution,
144            partitioning_to_restore: None,
145            preserve_hash_partitioning: false,
146        }
147    }
148}
149
150fn provided_global_fetch(plan: &Arc<dyn ExecutionPlan>) -> Option<usize> {
151    let fetch = plan.fetch()?;
152    (plan.as_any().is::<GlobalLimitExec>()
153        || plan.as_any().is::<CoalescePartitionsExec>()
154        || plan.as_any().is::<SortPreservingMergeExec>()
155        || plan.as_any().is::<MergeSortExec>())
156    .then_some(fetch)
157}
158
159fn add_global_fetch(
160    plan: Arc<dyn ExecutionPlan>,
161    fetch: usize,
162    required_ordering: Option<OrderingRequirements>,
163    partitioning_to_restore: Option<Partitioning>,
164) -> DfResult<Arc<dyn ExecutionPlan>> {
165    let plan = if required_ordering.is_some()
166        && let Some(ordering) = plan.output_ordering().cloned()
167    {
168        Arc::new(SortPreservingMergeExec::new(ordering, plan).with_fetch(Some(fetch)))
169            as Arc<dyn ExecutionPlan>
170    } else {
171        Arc::new(CoalescePartitionsExec::new(plan).with_fetch(Some(fetch)))
172            as Arc<dyn ExecutionPlan>
173    };
174
175    restore_required_partitioning(plan, partitioning_to_restore)
176}
177
178fn restore_required_partitioning(
179    plan: Arc<dyn ExecutionPlan>,
180    partitioning_to_restore: Option<Partitioning>,
181) -> DfResult<Arc<dyn ExecutionPlan>> {
182    let Some(partitioning) = partitioning_to_restore else {
183        return Ok(plan);
184    };
185
186    if partitioning.partition_count() <= 1 || !matches!(&partitioning, Partitioning::Hash(_, _)) {
187        return Ok(plan);
188    }
189
190    Ok(Arc::new(
191        RepartitionExec::try_new(plan, partitioning)?.with_preserve_order(),
192    ))
193}
194
195fn partitioning_to_restore_for(
196    child: &Arc<dyn ExecutionPlan>,
197    required_distribution: &Distribution,
198) -> Option<Partitioning> {
199    if !matches!(required_distribution, Distribution::HashPartitioned(_))
200        || child.output_partitioning().partition_count() <= 1
201    {
202        return None;
203    }
204
205    if child
206        .output_partitioning()
207        .satisfaction(required_distribution, child.equivalence_properties(), false)
208        .is_satisfied()
209    {
210        Some(child.output_partitioning().clone())
211    } else {
212        Some(
213            required_distribution
214                .clone()
215                .create_partitioning(child.output_partitioning().partition_count()),
216        )
217    }
218}
219
220fn inherited_partitioning_to_restore(
221    plan: &Arc<dyn ExecutionPlan>,
222    child: &Arc<dyn ExecutionPlan>,
223    parent: &ParentContext,
224) -> Option<Partitioning> {
225    if child.output_partitioning().partition_count() <= 1
226        || !matches!(child.output_partitioning(), Partitioning::Hash(_, _))
227        || !matches!(plan.output_partitioning(), Partitioning::Hash(_, _))
228        || plan.output_partitioning().partition_count()
229            != child.output_partitioning().partition_count()
230    {
231        return None;
232    }
233
234    let satisfies_parent_distribution = matches!(
235        parent.required_distribution,
236        Distribution::HashPartitioned(_)
237    ) && plan
238        .output_partitioning()
239        .satisfaction(
240            &parent.required_distribution,
241            plan.equivalence_properties(),
242            false,
243        )
244        .is_satisfied();
245
246    (satisfies_parent_distribution || parent.preserve_hash_partitioning)
247        .then(|| child.output_partitioning().clone())
248}
249
250#[cfg(test)]
251mod tests {
252    use datafusion::arrow::array::Int32Array;
253    use datafusion::arrow::compute::SortOptions;
254    use datafusion::arrow::datatypes::{DataType, Field, Schema};
255    use datafusion::arrow::record_batch::RecordBatch;
256    use datafusion::physical_expr::expressions::{col, lit};
257    use datafusion::physical_plan::filter::FilterExecBuilder;
258    use datafusion::physical_plan::joins::{HashJoinExec, PartitionMode};
259    use datafusion::physical_plan::limit::GlobalLimitExec;
260    use datafusion::physical_plan::projection::ProjectionExec;
261    use datafusion::physical_plan::repartition::RepartitionExec;
262    use datafusion::physical_plan::test::TestMemoryExec;
263    use datafusion_common::{JoinType, NullEquality};
264    use datafusion_physical_expr::{LexOrdering, Partitioning, PhysicalSortExpr};
265
266    use super::*;
267
268    #[test]
269    fn adds_global_limit_for_multi_partition_filter_fetch() {
270        let filter = filter_fetch(unordered_input(), 1);
271
272        let optimized =
273            EnsureGlobalLimitForFetch::optimize_plan(filter, ParentContext::default()).unwrap();
274
275        assert!(optimized.as_any().is::<CoalescePartitionsExec>());
276        assert_eq!(optimized.fetch(), Some(1));
277        assert_eq!(optimized.output_partitioning().partition_count(), 1);
278    }
279
280    #[test]
281    fn still_visits_subtree_under_global_limit() {
282        let filter = filter_fetch(unordered_input(), 5);
283        let projection = Arc::new(
284            ProjectionExec::try_new(
285                vec![(col("a", filter.schema().as_ref()).unwrap(), "a".to_string())],
286                filter,
287            )
288            .unwrap(),
289        );
290        let limit =
291            Arc::new(GlobalLimitExec::new(projection, 0, Some(10))) as Arc<dyn ExecutionPlan>;
292
293        let optimized =
294            EnsureGlobalLimitForFetch::optimize_plan(limit, ParentContext::default()).unwrap();
295        let projection = optimized.children()[0];
296        let coalesce = projection.children()[0];
297
298        assert!(coalesce.as_any().is::<CoalescePartitionsExec>());
299        assert_eq!(coalesce.fetch(), Some(5));
300    }
301
302    #[test]
303    fn keeps_filter_under_parent_global_fetch() {
304        let (input, ordering) = ordered_input();
305        let filter = filter_fetch(input, 1);
306        let merge = Arc::new(SortPreservingMergeExec::new(ordering, filter).with_fetch(Some(1)))
307            as Arc<dyn ExecutionPlan>;
308
309        let optimized =
310            EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
311        let child = optimized.children()[0];
312
313        assert!(optimized.as_any().is::<SortPreservingMergeExec>());
314        assert!(child.as_any().is::<FilterExec>());
315    }
316
317    #[test]
318    fn adds_tighter_global_fetch_under_looser_parent_fetch() {
319        let (input, ordering) = ordered_input();
320        let filter = filter_fetch(input, 5);
321        let merge = Arc::new(SortPreservingMergeExec::new(ordering, filter).with_fetch(Some(10)))
322            as Arc<dyn ExecutionPlan>;
323
324        let optimized =
325            EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
326        let child = optimized.children()[0];
327
328        assert!(optimized.as_any().is::<SortPreservingMergeExec>());
329        assert!(child.as_any().is::<SortPreservingMergeExec>());
330        assert_eq!(child.fetch(), Some(5));
331    }
332
333    #[test]
334    fn keeps_filter_under_parent_merge_sort_fetch() {
335        let (input, ordering) = ordered_input();
336        let filter = filter_fetch(input, 1);
337        let merge = merge_sort_fetch(ordering, filter, 1);
338
339        let optimized =
340            EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
341        let child = optimized.children()[0];
342
343        assert!(optimized.as_any().is::<MergeSortExec>());
344        assert!(child.as_any().is::<FilterExec>());
345    }
346
347    #[test]
348    fn adds_tighter_global_fetch_under_looser_merge_sort_fetch() {
349        let (input, ordering) = ordered_input();
350        let filter = filter_fetch(input, 5);
351        let merge = merge_sort_fetch(ordering, filter, 10);
352
353        let optimized =
354            EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
355        let child = optimized.children()[0];
356
357        assert!(optimized.as_any().is::<MergeSortExec>());
358        assert!(child.as_any().is::<SortPreservingMergeExec>());
359        assert_eq!(child.fetch(), Some(5));
360        assert!(child.children()[0].as_any().is::<FilterExec>());
361    }
362
363    #[test]
364    fn preserves_parent_ordering_requirement() {
365        let (input, ordering) = ordered_input();
366        let filter = filter_fetch(input, 1);
367        let merge =
368            Arc::new(SortPreservingMergeExec::new(ordering, filter)) as Arc<dyn ExecutionPlan>;
369
370        let optimized =
371            EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
372        let child = optimized.children()[0];
373
374        assert!(optimized.as_any().is::<SortPreservingMergeExec>());
375        assert!(child.as_any().is::<SortPreservingMergeExec>());
376        assert_eq!(child.fetch(), Some(1));
377    }
378
379    #[test]
380    fn uses_child_output_ordering_for_merge() {
381        let schema = schema();
382        let required_ordering = ordering(schema.as_ref(), false);
383        let actual_ordering = ordering(schema.as_ref(), true);
384        let batch = batch(schema.clone());
385        let partitions = vec![vec![batch.clone()], vec![batch.clone()], vec![batch]];
386        let input = TestMemoryExec::try_new(&partitions, schema, None)
387            .unwrap()
388            .try_with_sort_information(vec![actual_ordering.clone()])
389            .unwrap();
390        let filter = filter_fetch(Arc::new(input), 1);
391
392        let optimized = add_global_fetch(
393            filter,
394            1,
395            Some(OrderingRequirements::from(required_ordering)),
396            None,
397        )
398        .unwrap();
399        let merge = optimized
400            .as_any()
401            .downcast_ref::<SortPreservingMergeExec>()
402            .unwrap();
403
404        assert_eq!(merge.expr(), &actual_ordering);
405    }
406
407    #[test]
408    fn preserves_inherited_ordering_requirement_through_projection() {
409        let (input, ordering) = ordered_input();
410        let filter = filter_fetch(input, 1);
411        let projection = Arc::new(
412            ProjectionExec::try_new(
413                vec![(col("a", filter.schema().as_ref()).unwrap(), "a".to_string())],
414                filter,
415            )
416            .unwrap(),
417        );
418        let merge =
419            Arc::new(SortPreservingMergeExec::new(ordering, projection)) as Arc<dyn ExecutionPlan>;
420
421        let optimized =
422            EnsureGlobalLimitForFetch::optimize_plan(merge, ParentContext::default()).unwrap();
423        let projection = optimized.children()[0];
424        let child = projection.children()[0];
425
426        assert!(optimized.as_any().is::<SortPreservingMergeExec>());
427        assert!(projection.as_any().is::<ProjectionExec>());
428        assert!(child.as_any().is::<SortPreservingMergeExec>());
429        assert_eq!(child.fetch(), Some(1));
430    }
431
432    #[test]
433    fn restores_parent_hash_distribution_after_global_fetch() {
434        let left = filter_fetch(hash_repartition(unordered_input()), 1);
435        let right = hash_repartition(unordered_input());
436        let on = vec![(
437            col("a", left.schema().as_ref()).unwrap(),
438            col("a", right.schema().as_ref()).unwrap(),
439        )];
440        let join = Arc::new(
441            HashJoinExec::try_new(
442                left,
443                right,
444                on,
445                None,
446                &JoinType::Inner,
447                None,
448                PartitionMode::Partitioned,
449                NullEquality::NullEqualsNothing,
450                false,
451            )
452            .unwrap(),
453        ) as Arc<dyn ExecutionPlan>;
454
455        let optimized =
456            EnsureGlobalLimitForFetch::optimize_plan(join, ParentContext::default()).unwrap();
457        let left = optimized.children()[0];
458        let repartition = left.as_any().downcast_ref::<RepartitionExec>().unwrap();
459
460        assert!(matches!(
461            repartition.partitioning(),
462            Partitioning::Hash(_, 3)
463        ));
464        assert!(repartition.input().as_any().is::<CoalescePartitionsExec>());
465        assert_eq!(repartition.input().fetch(), Some(1));
466    }
467
468    #[test]
469    fn restores_inherited_hash_distribution_through_projection() {
470        let filter = filter_fetch(hash_repartition(unordered_input()), 1);
471        let projection = Arc::new(
472            ProjectionExec::try_new(
473                vec![(col("a", filter.schema().as_ref()).unwrap(), "a".to_string())],
474                filter,
475            )
476            .unwrap(),
477        ) as Arc<dyn ExecutionPlan>;
478        let right = hash_repartition(unordered_input());
479        let on = vec![(
480            col("a", projection.schema().as_ref()).unwrap(),
481            col("a", right.schema().as_ref()).unwrap(),
482        )];
483        let join = Arc::new(
484            HashJoinExec::try_new(
485                projection,
486                right,
487                on,
488                None,
489                &JoinType::Inner,
490                None,
491                PartitionMode::Partitioned,
492                NullEquality::NullEqualsNothing,
493                false,
494            )
495            .unwrap(),
496        ) as Arc<dyn ExecutionPlan>;
497
498        let optimized =
499            EnsureGlobalLimitForFetch::optimize_plan(join, ParentContext::default()).unwrap();
500        let projection = optimized.children()[0];
501        let repartition = projection.children()[0]
502            .as_any()
503            .downcast_ref::<RepartitionExec>()
504            .unwrap();
505
506        assert!(projection.as_any().is::<ProjectionExec>());
507        assert!(matches!(
508            repartition.partitioning(),
509            Partitioning::Hash(_, 3)
510        ));
511        assert!(repartition.input().as_any().is::<CoalescePartitionsExec>());
512        assert_eq!(repartition.input().fetch(), Some(1));
513    }
514
515    #[test]
516    fn restores_inherited_hash_distribution_through_multiple_projections() {
517        let filter = filter_fetch(hash_repartition(unordered_input()), 1);
518        let projection = project_a(filter);
519        let projection = project_a(projection);
520        let right = hash_repartition(unordered_input());
521        let on = vec![(
522            col("a", projection.schema().as_ref()).unwrap(),
523            col("a", right.schema().as_ref()).unwrap(),
524        )];
525        let join = Arc::new(
526            HashJoinExec::try_new(
527                projection,
528                right,
529                on,
530                None,
531                &JoinType::Inner,
532                None,
533                PartitionMode::Partitioned,
534                NullEquality::NullEqualsNothing,
535                false,
536            )
537            .unwrap(),
538        ) as Arc<dyn ExecutionPlan>;
539
540        let optimized =
541            EnsureGlobalLimitForFetch::optimize_plan(join, ParentContext::default()).unwrap();
542        let outer_projection = optimized.children()[0];
543        let inner_projection = outer_projection.children()[0];
544        let repartition = inner_projection.children()[0]
545            .as_any()
546            .downcast_ref::<RepartitionExec>()
547            .unwrap();
548
549        assert!(outer_projection.as_any().is::<ProjectionExec>());
550        assert!(inner_projection.as_any().is::<ProjectionExec>());
551        assert!(matches!(
552            repartition.partitioning(),
553            Partitioning::Hash(_, 3)
554        ));
555        assert!(repartition.input().as_any().is::<CoalescePartitionsExec>());
556        assert_eq!(repartition.input().fetch(), Some(1));
557    }
558
559    fn unordered_input() -> Arc<dyn ExecutionPlan> {
560        let schema = schema();
561        let batch = batch(schema.clone());
562        let partitions = vec![vec![batch.clone()], vec![batch.clone()], vec![batch]];
563        Arc::new(TestMemoryExec::try_new(&partitions, schema, None).unwrap())
564    }
565
566    fn ordered_input() -> (Arc<dyn ExecutionPlan>, LexOrdering) {
567        let schema = schema();
568        let ordering = ordering(schema.as_ref(), false);
569        let batch = batch(schema.clone());
570        let partitions = vec![vec![batch.clone()], vec![batch.clone()], vec![batch]];
571        let input = TestMemoryExec::try_new(&partitions, schema, None)
572            .unwrap()
573            .try_with_sort_information(vec![ordering.clone()])
574            .unwrap();
575
576        (Arc::new(input), ordering)
577    }
578
579    fn filter_fetch(input: Arc<dyn ExecutionPlan>, fetch: usize) -> Arc<dyn ExecutionPlan> {
580        Arc::new(
581            FilterExecBuilder::new(lit(true), input)
582                .with_fetch(Some(fetch))
583                .build()
584                .unwrap(),
585        )
586    }
587
588    fn merge_sort_fetch(
589        ordering: LexOrdering,
590        input: Arc<dyn ExecutionPlan>,
591        fetch: usize,
592    ) -> Arc<dyn ExecutionPlan> {
593        Arc::new(MergeSortExec::new(ordering, input, Some(fetch)))
594    }
595
596    fn hash_repartition(input: Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
597        let partitioning = Partitioning::Hash(vec![col("a", input.schema().as_ref()).unwrap()], 3);
598        Arc::new(RepartitionExec::try_new(input, partitioning).unwrap())
599    }
600
601    fn project_a(input: Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
602        Arc::new(
603            ProjectionExec::try_new(
604                vec![(col("a", input.schema().as_ref()).unwrap(), "a".to_string())],
605                input,
606            )
607            .unwrap(),
608        )
609    }
610
611    fn schema() -> Arc<Schema> {
612        Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)]))
613    }
614
615    fn batch(schema: Arc<Schema>) -> RecordBatch {
616        RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 2, 3]))]).unwrap()
617    }
618
619    fn ordering(schema: &Schema, descending: bool) -> LexOrdering {
620        LexOrdering::new([PhysicalSortExpr::new(
621            col("a", schema).unwrap(),
622            SortOptions {
623                descending,
624                nulls_first: descending,
625            },
626        )])
627        .unwrap()
628    }
629}