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::{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}