Skip to main content

query/dist_plan/
merge_sort.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
15//! Merge sort logical plan for distributed query execution, roughly corresponding to the
16//! `SortPreservingMergeExec` operator in datafusion
17//!
18
19use std::fmt;
20use std::sync::Arc;
21
22use datafusion::execution::TaskContext;
23use datafusion::physical_plan::execution_plan::CardinalityEffect;
24use datafusion::physical_plan::metrics::MetricsSet;
25use datafusion::physical_plan::projection::{ProjectionExec, make_with_child, update_ordering};
26use datafusion::physical_plan::sorts::sort::SortExec;
27use datafusion::physical_plan::sorts::sort_preserving_merge::SortPreservingMergeExec;
28use datafusion::physical_plan::{
29    ChildStats, ChildrenPropertiesMode, DisplayAs, DisplayFormatType, ExecutionPlan,
30    InputDistributionRequirements, PlanProperties, ReplaceChildrenOptions,
31    SendableRecordBatchStream, Statistics, StatisticsArgs, apply_expression_roots,
32};
33use datafusion_common::tree_node::TreeNodeRecursion;
34use datafusion_common::{DataFusionError, Result};
35use datafusion_expr::{Extension, LogicalPlan, SortExpr, UserDefinedLogicalNodeCore};
36use datafusion_physical_expr::{LexOrdering, OrderingRequirements, PhysicalExpr};
37
38/// MergeSort Logical Plan, have same field as `Sort`, but indicate it is a merge sort,
39/// which assume each input partition is a sorted stream, and will use `SortPreserveingMergeExec`
40/// to merge them into a single sorted stream.
41#[derive(Hash, PartialOrd, PartialEq, Eq, Clone)]
42pub struct MergeSortLogicalPlan {
43    pub expr: Vec<SortExpr>,
44    pub input: Arc<LogicalPlan>,
45    pub fetch: Option<usize>,
46}
47
48impl fmt::Debug for MergeSortLogicalPlan {
49    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
50        UserDefinedLogicalNodeCore::fmt_for_explain(self, f)
51    }
52}
53
54impl MergeSortLogicalPlan {
55    pub fn new(input: Arc<LogicalPlan>, expr: Vec<SortExpr>, fetch: Option<usize>) -> Self {
56        Self { input, expr, fetch }
57    }
58
59    pub fn name() -> &'static str {
60        "MergeSort"
61    }
62
63    /// Create a [`LogicalPlan::Extension`] node from this merge sort plan
64    pub fn into_logical_plan(self) -> LogicalPlan {
65        LogicalPlan::Extension(Extension {
66            node: Arc::new(self),
67        })
68    }
69}
70
71/// An opaque physical execution node for [`MergeSortLogicalPlan`].
72///
73/// It delegates execution and physical properties to DataFusion's
74/// [`SortPreservingMergeExec`], but intentionally does not expose itself as a
75/// `SortPreservingMergeExec`. `EnforceSorting` is allowed to replace a bare
76/// `SortPreservingMergeExec` with `CoalescePartitionsExec` when the parent does
77/// not require ordering. `MergeSortExec` represents the distributed TopK merge
78/// stage itself, so later physical optimizer rules must not rewrite it into an
79/// unordered fetch.
80#[derive(Debug, Clone)]
81pub(crate) struct MergeSortExec {
82    inner: SortPreservingMergeExec,
83}
84
85impl MergeSortExec {
86    pub(crate) fn new(
87        ordering: LexOrdering,
88        input: Arc<dyn ExecutionPlan>,
89        fetch: Option<usize>,
90    ) -> Self {
91        Self {
92            inner: SortPreservingMergeExec::new(ordering, input).with_fetch(fetch),
93        }
94    }
95
96    fn input_with_fetch(&self, fetch: Option<usize>) -> Arc<dyn ExecutionPlan> {
97        let input = Arc::clone(self.inner.input());
98        if let Some(sort) = input.downcast_ref::<SortExec>()
99            && sort.preserve_partitioning()
100            && sort.expr() == self.inner.expr()
101        {
102            // Mirror DataFusion's bare SPM plan quality for distributed TopK:
103            // keep the parent `MergeSortExec(fetch)` as the global merge, and
104            // bound the partition-preserving child sort to the same local TopK.
105            // Local top-K is safe because every global top-K row must be within
106            // the top-K rows of its own input partition.
107            Arc::new(sort.with_fetch(fetch))
108        } else {
109            input
110        }
111    }
112}
113
114impl DisplayAs for MergeSortExec {
115    fn fmt_as(&self, t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result {
116        match t {
117            DisplayFormatType::Default | DisplayFormatType::Verbose => {
118                write!(f, "MergeSortExec: [{}]", self.inner.expr())?;
119                if let Some(fetch) = self.inner.fetch() {
120                    write!(f, ", fetch={fetch}")?;
121                }
122                Ok(())
123            }
124            DisplayFormatType::TreeRender => {
125                if let Some(fetch) = self.inner.fetch() {
126                    writeln!(f, "limit={fetch}")?;
127                }
128
129                for (i, expr) in self.inner.expr().iter().enumerate() {
130                    expr.fmt_sql(f)?;
131                    if i != self.inner.expr().len() - 1 {
132                        write!(f, ", ")?;
133                    }
134                }
135
136                Ok(())
137            }
138        }
139    }
140}
141
142impl ExecutionPlan for MergeSortExec {
143    fn name(&self) -> &str {
144        "MergeSortExec"
145    }
146
147    /// Keeps this node intentionally opaque to DataFusion's type-specialized
148    /// optimizer rewrites.
149    ///
150    /// `MergeSortExec` delegates most behavior to DataFusion's
151    /// `SortPreservingMergeExec`, but it must not expose itself as that type.
152    /// DataFusion's `EnforceSorting` optimizer recognizes a bare
153    /// `SortPreservingMergeExec` via `downcast_ref::<...>()` and may
154    /// replace it with an unordered `CoalescePartitionsExec(fetch)` when the
155    /// parent does not require sorted output.
156    ///
157    /// That rewrite is valid for an ordinary SPM used only to satisfy parent
158    /// ordering, but not for GreptimeDB's distributed TopK merge stage. In a
159    /// scalar-subquery shape like `ORDER BY ts DESC LIMIT 1`, this node is the
160    /// operator that merges region-local TopK streams into the global TopK.
161    /// Replacing it with unordered coalescing can return a partial/latest row
162    /// from one region instead of the global latest row.
163    ///
164    /// `required_input_ordering()` separately tells DataFusion what ordering this
165    /// node needs from its child, so `EnforceSorting` can insert a `SortExec`
166    /// below `MergeSortExec` when `MergeScanExec` cannot preserve per-partition
167    /// ordering. This opacity is specifically about protecting the merge stage
168    /// itself from the `EnforceSorting` rewrite above.
169    fn properties(&self) -> &Arc<PlanProperties> {
170        self.inner.properties()
171    }
172
173    /// Forwards DataFusion's order-preserving scan hint through this wrapper.
174    ///
175    /// This mirrors `SortPreservingMergeExec::with_preserve_order()`: if the
176    /// child can produce an order-preserving variant, rebuild the same merge
177    /// stage on top of that child. The returned plan must stay a
178    /// `MergeSortExec`, not a bare SPM, so the distributed TopK merge remains
179    /// opaque to `EnforceSorting`'s SPM-specific rewrite.
180    fn with_preserve_order(&self, preserve_order: bool) -> Option<Arc<dyn ExecutionPlan>> {
181        self.inner
182            .input()
183            .with_preserve_order(preserve_order)
184            .map(|new_input| {
185                Arc::new(Self::new(
186                    self.inner.expr().clone(),
187                    new_input,
188                    self.inner.fetch(),
189                )) as Arc<dyn ExecutionPlan>
190            })
191    }
192
193    fn input_distribution_requirements(&self) -> InputDistributionRequirements {
194        self.inner.input_distribution_requirements()
195    }
196
197    fn benefits_from_input_partitioning(&self) -> Vec<bool> {
198        self.inner.benefits_from_input_partitioning()
199    }
200
201    /// Tells DataFusion that `MergeSortExec` requires each input partition to be
202    /// ordered. This is the contract that makes `EnforceSorting` insert a
203    /// `SortExec` below `MergeSortExec` when the input cannot preserve ordering.
204    ///
205    /// The opacity of `MergeSortExec`'s downcast identity, not this requirement, is what
206    /// prevents DataFusion from rewriting the merge stage itself as a bare
207    /// `SortPreservingMergeExec`.
208    fn required_input_ordering(&self) -> Vec<Option<OrderingRequirements>> {
209        vec![Some(OrderingRequirements::from(self.inner.expr().clone()))]
210    }
211
212    fn maintains_input_order(&self) -> Vec<bool> {
213        self.inner.maintains_input_order()
214    }
215
216    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
217        self.inner.children()
218    }
219
220    fn apply_expressions(
221        &self,
222        f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> Result<TreeNodeRecursion>,
223    ) -> Result<TreeNodeRecursion> {
224        apply_expression_roots(self.inner.expr().iter().map(|sort_expr| &sort_expr.expr), f)
225    }
226
227    fn replace_children(
228        self: Arc<Self>,
229        mut children: Vec<Arc<dyn ExecutionPlan>>,
230        options: ReplaceChildrenOptions,
231    ) -> Result<Arc<dyn ExecutionPlan>> {
232        if children.len() != 1 {
233            return Err(DataFusionError::Internal(format!(
234                "MergeSortExec expects exactly one child, got {}",
235                children.len()
236            )));
237        }
238
239        match options.children_properties {
240            ChildrenPropertiesMode::Keep => Ok(Arc::new(Self {
241                inner: SortPreservingMergeExec::new(
242                    self.inner.expr().clone(),
243                    children.swap_remove(0),
244                )
245                .with_fetch(self.inner.fetch()),
246            })),
247            ChildrenPropertiesMode::Recompute => Ok(Arc::new(Self::new(
248                self.inner.expr().clone(),
249                children.swap_remove(0),
250                self.inner.fetch(),
251            ))),
252        }
253    }
254
255    #[allow(deprecated)]
256    fn with_new_children(
257        self: Arc<Self>,
258        children: Vec<Arc<dyn ExecutionPlan>>,
259    ) -> Result<Arc<dyn ExecutionPlan>> {
260        self.replace_children(
261            children,
262            ReplaceChildrenOptions::new(ChildrenPropertiesMode::Recompute),
263        )
264    }
265
266    fn execute(
267        &self,
268        partition: usize,
269        context: Arc<TaskContext>,
270    ) -> Result<SendableRecordBatchStream> {
271        self.inner.execute(partition, context)
272    }
273
274    fn metrics(&self) -> Option<MetricsSet> {
275        self.inner.metrics()
276    }
277
278    fn child_stats_requests(&self, partition: Option<usize>) -> Vec<ChildStats> {
279        self.inner.child_stats_requests(partition)
280    }
281
282    fn statistics_from_inputs(
283        &self,
284        input_stats: &[Arc<Statistics>],
285        args: &StatisticsArgs,
286    ) -> Result<Arc<Statistics>> {
287        self.inner.statistics_from_inputs(input_stats, args)
288    }
289
290    fn cardinality_effect(&self) -> CardinalityEffect {
291        self.inner.cardinality_effect()
292    }
293
294    /// Intentionally keeps DataFusion's generic limit pushdown disabled.
295    ///
296    /// `MergeSortExec` still supports its own global fetch through
297    /// `with_fetch()`. What we must not allow is pushing an external limit below
298    /// this required distributed TopK merge. DataFusion's limit pushdown rules
299    /// know how to treat a bare `SortPreservingMergeExec` as a
300    /// partition-combining node, but `MergeSortExec` is intentionally opaque to
301    /// those SPM-specific downcasts. Enabling generic limit pushdown without also
302    /// teaching the optimizer about this wrapper could return partition-local
303    /// rows instead of the global TopK.
304    fn supports_limit_pushdown(&self) -> bool {
305        false
306    }
307
308    fn fetch(&self) -> Option<usize> {
309        self.inner.fetch()
310    }
311
312    fn with_fetch(&self, limit: Option<usize>) -> Option<Arc<dyn ExecutionPlan>> {
313        Some(Arc::new(Self::new(
314            self.inner.expr().clone(),
315            self.input_with_fetch(limit),
316            limit,
317        )))
318    }
319
320    /// Lets DataFusion push a projection below this merge when it can rewrite
321    /// the ordering expressions safely.
322    ///
323    /// This mirrors `SortPreservingMergeExec::try_swapping_with_projection()`
324    /// for plan quality, but re-wraps the result as `MergeSortExec` so the
325    /// distributed merge stage keeps its type identity and opacity.
326    fn try_swapping_with_projection(
327        &self,
328        projection: &ProjectionExec,
329    ) -> Result<Option<Arc<dyn ExecutionPlan>>> {
330        if projection.expr().len() >= projection.input().schema().fields().len() {
331            return Ok(None);
332        }
333
334        let Some(updated_exprs) = update_ordering(self.inner.expr().clone(), projection.expr())?
335        else {
336            return Ok(None);
337        };
338
339        Ok(Some(Arc::new(Self::new(
340            updated_exprs,
341            make_with_child(projection, self.inner.input())?,
342            self.inner.fetch(),
343        ))))
344    }
345}
346
347impl UserDefinedLogicalNodeCore for MergeSortLogicalPlan {
348    fn name(&self) -> &str {
349        Self::name()
350    }
351
352    // Allow optimization here
353    fn inputs(&self) -> Vec<&LogicalPlan> {
354        vec![self.input.as_ref()]
355    }
356
357    fn schema(&self) -> &datafusion_common::DFSchemaRef {
358        self.input.schema()
359    }
360
361    // Allow further optimization
362    fn expressions(&self) -> Vec<datafusion_expr::Expr> {
363        self.expr.iter().map(|sort| sort.expr.clone()).collect()
364    }
365
366    fn fmt_for_explain(&self, f: &mut fmt::Formatter) -> fmt::Result {
367        write!(f, "MergeSort: ")?;
368        for (i, expr_item) in self.expr.iter().enumerate() {
369            if i > 0 {
370                write!(f, ", ")?;
371            }
372            write!(f, "{expr_item}")?;
373        }
374        if let Some(a) = self.fetch {
375            write!(f, ", fetch={a}")?;
376        }
377        Ok(())
378    }
379
380    fn with_exprs_and_inputs(
381        &self,
382        exprs: Vec<datafusion::prelude::Expr>,
383        mut inputs: Vec<LogicalPlan>,
384    ) -> Result<Self> {
385        let mut zelf = self.clone();
386        zelf.expr = zelf
387            .expr
388            .into_iter()
389            .zip(exprs)
390            .map(|(sort, expr)| sort.with_expr(expr))
391            .collect();
392        zelf.input = Arc::new(inputs.pop().ok_or_else(|| {
393            DataFusionError::Internal("Expected exactly one input with MergeSort".to_string())
394        })?);
395        Ok(zelf)
396    }
397}
398
399/// Turn `Sort` into `MergeSort` if possible
400pub fn merge_sort_transformer(plan: &LogicalPlan) -> Option<LogicalPlan> {
401    if let LogicalPlan::Sort(sort) = plan {
402        Some(
403            MergeSortLogicalPlan::new(sort.input.clone(), sort.expr.clone(), sort.fetch)
404                .into_logical_plan(),
405        )
406    } else {
407        None
408    }
409}
410
411#[cfg(test)]
412mod tests {
413    use arrow_schema::{DataType, Field, Schema, SortOptions};
414    use datafusion::physical_optimizer::enforce_sorting::replace_with_order_preserving_variants::{
415        OrderPreservationContext, plan_with_order_breaking_variants,
416    };
417    use datafusion::physical_plan::coalesce_partitions::CoalescePartitionsExec;
418    use datafusion::physical_plan::displayable;
419    use datafusion::physical_plan::empty::EmptyExec;
420    use datafusion_physical_expr::PhysicalSortExpr;
421    use datafusion_physical_expr::expressions::col as physical_col;
422
423    use super::*;
424
425    /// Test double that records DataFusion's preserve-order signal while
426    /// otherwise behaving like a transparent wrapper around its child.
427    #[derive(Debug, Clone)]
428    struct PreserveOrderProbeExec {
429        inner: Arc<dyn ExecutionPlan>,
430        preserve_order: bool,
431    }
432
433    impl PreserveOrderProbeExec {
434        fn new(inner: Arc<dyn ExecutionPlan>) -> Self {
435            Self {
436                inner,
437                preserve_order: false,
438            }
439        }
440    }
441
442    impl DisplayAs for PreserveOrderProbeExec {
443        fn fmt_as(&self, _t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result {
444            write!(
445                f,
446                "PreserveOrderProbeExec: preserve_order={}",
447                self.preserve_order
448            )
449        }
450    }
451
452    impl ExecutionPlan for PreserveOrderProbeExec {
453        fn name(&self) -> &str {
454            "PreserveOrderProbeExec"
455        }
456
457        fn properties(&self) -> &Arc<PlanProperties> {
458            self.inner.properties()
459        }
460
461        fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
462            vec![&self.inner]
463        }
464
465        fn apply_expressions(
466            &self,
467            _f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> Result<TreeNodeRecursion>,
468        ) -> Result<TreeNodeRecursion> {
469            Ok(TreeNodeRecursion::Continue)
470        }
471
472        fn with_new_children(
473            self: Arc<Self>,
474            mut children: Vec<Arc<dyn ExecutionPlan>>,
475        ) -> Result<Arc<dyn ExecutionPlan>> {
476            if children.len() != 1 {
477                return Err(DataFusionError::Internal(format!(
478                    "PreserveOrderProbeExec expects exactly one child, got {}",
479                    children.len()
480                )));
481            }
482
483            Ok(Arc::new(Self {
484                inner: children.swap_remove(0),
485                preserve_order: self.preserve_order,
486            }))
487        }
488
489        fn execute(
490            &self,
491            partition: usize,
492            context: Arc<TaskContext>,
493        ) -> Result<SendableRecordBatchStream> {
494            self.inner.execute(partition, context)
495        }
496
497        fn with_preserve_order(&self, preserve_order: bool) -> Option<Arc<dyn ExecutionPlan>> {
498            Some(Arc::new(Self {
499                inner: Arc::clone(&self.inner),
500                preserve_order,
501            }))
502        }
503    }
504
505    fn test_ordering(schema: &Schema) -> LexOrdering {
506        LexOrdering::new([PhysicalSortExpr::new(
507            physical_col("ts", schema).unwrap(),
508            SortOptions {
509                descending: true,
510                nulls_first: false,
511            },
512        )])
513        .unwrap()
514    }
515
516    #[test]
517    fn merge_sort_exec_is_opaque_and_preserves_topk_requirements() {
518        let schema = Arc::new(Schema::new(vec![Field::new("ts", DataType::Int64, false)]));
519        let input = Arc::new(EmptyExec::new(schema.clone()).with_partitions(2)) as _;
520        let ordering = test_ordering(schema.as_ref());
521
522        let merge_sort =
523            Arc::new(MergeSortExec::new(ordering, input, Some(1))) as Arc<dyn ExecutionPlan>;
524
525        assert_eq!(merge_sort.name(), "MergeSortExec");
526        assert!(
527            merge_sort
528                .downcast_ref::<SortPreservingMergeExec>()
529                .is_none(),
530            "MergeSortExec must stay opaque to EnforceSorting's bare SortPreservingMerge rewrite"
531        );
532        assert_eq!(merge_sort.fetch(), Some(1));
533        assert!(!merge_sort.supports_limit_pushdown());
534        assert!(merge_sort.required_input_ordering()[0].is_some());
535
536        let tree = displayable(merge_sort.as_ref()).tree_render().to_string();
537        assert!(tree.contains("MergeSortExec"));
538        assert!(!tree.contains("SortPreservingMergeExec"));
539
540        let fetched = merge_sort.with_fetch(Some(2)).unwrap();
541        assert!(fetched.downcast_ref::<MergeSortExec>().is_some());
542        assert_eq!(fetched.fetch(), Some(2));
543    }
544
545    #[test]
546    fn merge_sort_exec_required_input_ordering_matches_spm() {
547        let schema = Arc::new(Schema::new(vec![Field::new("ts", DataType::Int64, false)]));
548        let input = Arc::new(EmptyExec::new(schema.clone()).with_partitions(2)) as _;
549        let ordering = test_ordering(schema.as_ref());
550
551        let merge_sort = MergeSortExec::new(ordering.clone(), Arc::clone(&input), Some(1));
552        let bare_spm =
553            SortPreservingMergeExec::new(ordering.clone(), Arc::clone(&input)).with_fetch(Some(1));
554
555        assert_eq!(
556            merge_sort.required_input_ordering(),
557            vec![Some(OrderingRequirements::from(ordering))],
558            "MergeSortExec must require locally sorted input partitions for the merge key"
559        );
560        assert_eq!(
561            merge_sort.required_input_ordering(),
562            bare_spm.required_input_ordering(),
563            "MergeSortExec's child ordering contract should mirror SortPreservingMergeExec"
564        );
565        assert_eq!(
566            merge_sort.maintains_input_order(),
567            bare_spm.maintains_input_order()
568        );
569    }
570
571    #[test]
572    fn merge_sort_exec_with_fetch_pushes_fetch_to_child_sort() {
573        let schema = Arc::new(Schema::new(vec![Field::new("ts", DataType::Int64, false)]));
574        let input = Arc::new(EmptyExec::new(schema.clone()).with_partitions(2)) as _;
575        let ordering = test_ordering(schema.as_ref());
576        let child_sort =
577            Arc::new(SortExec::new(ordering.clone(), input).with_preserve_partitioning(true))
578                as Arc<dyn ExecutionPlan>;
579        let merge_sort = MergeSortExec::new(ordering, child_sort, None);
580
581        let fetched = merge_sort.with_fetch(Some(2)).unwrap();
582
583        assert!(fetched.downcast_ref::<MergeSortExec>().is_some());
584        assert_eq!(fetched.fetch(), Some(2));
585        let child_sort = fetched.children()[0].downcast_ref::<SortExec>().unwrap();
586        assert_eq!(child_sort.fetch(), Some(2));
587        assert!(child_sort.preserve_partitioning());
588    }
589
590    #[test]
591    fn merge_sort_exec_with_preserve_order_matches_spm_but_keeps_wrapper() {
592        let schema = Arc::new(Schema::new(vec![Field::new("ts", DataType::Int64, false)]));
593        let input = Arc::new(PreserveOrderProbeExec::new(Arc::new(
594            EmptyExec::new(schema.clone()).with_partitions(2),
595        ))) as _;
596        let ordering = test_ordering(schema.as_ref());
597
598        let bare_spm =
599            SortPreservingMergeExec::new(ordering.clone(), Arc::clone(&input)).with_fetch(Some(1));
600        let preserved_spm = bare_spm.with_preserve_order(true).unwrap();
601        assert!(
602            preserved_spm
603                .downcast_ref::<SortPreservingMergeExec>()
604                .is_some(),
605            "bare SPM should rebuild as bare SPM"
606        );
607        assert!(
608            preserved_spm.children()[0]
609                .downcast_ref::<PreserveOrderProbeExec>()
610                .unwrap()
611                .preserve_order
612        );
613
614        let merge_sort = MergeSortExec::new(ordering, input, Some(1));
615        let preserved_merge_sort = merge_sort.with_preserve_order(true).unwrap();
616        assert!(
617            preserved_merge_sort
618                .downcast_ref::<MergeSortExec>()
619                .is_some(),
620            "MergeSortExec must rewrap the preserve-order child as MergeSortExec"
621        );
622        assert!(
623            preserved_merge_sort
624                .downcast_ref::<SortPreservingMergeExec>()
625                .is_none(),
626            "MergeSortExec must not expose a bare SPM after with_preserve_order"
627        );
628        assert_eq!(preserved_merge_sort.fetch(), Some(1));
629        assert_eq!(
630            preserved_merge_sort.required_input_ordering(),
631            preserved_spm.required_input_ordering(),
632            "preserve-order rewrite should keep the same SPM child-ordering contract"
633        );
634        assert!(
635            preserved_merge_sort.children()[0]
636                .downcast_ref::<PreserveOrderProbeExec>()
637                .unwrap()
638                .preserve_order
639        );
640    }
641
642    #[test]
643    fn merge_sort_exec_projection_swap_matches_spm_but_keeps_wrapper() -> Result<()> {
644        let schema = Arc::new(Schema::new(vec![
645            Field::new("value", DataType::Int64, false),
646            Field::new("ts", DataType::Int64, false),
647            Field::new("tag", DataType::Utf8, false),
648        ]));
649        let input = Arc::new(EmptyExec::new(schema.clone()).with_partitions(2)) as _;
650        let ordering = test_ordering(schema.as_ref());
651
652        let bare_spm = Arc::new(
653            SortPreservingMergeExec::new(ordering.clone(), Arc::clone(&input)).with_fetch(Some(1)),
654        ) as Arc<dyn ExecutionPlan>;
655        let spm_projection = ProjectionExec::try_new(
656            vec![
657                (physical_col("ts", schema.as_ref())?, "ts".to_string()),
658                (physical_col("tag", schema.as_ref())?, "tag".to_string()),
659            ],
660            Arc::clone(&bare_spm),
661        )?;
662        let swapped_spm = bare_spm
663            .try_swapping_with_projection(&spm_projection)?
664            .expect("SPM should accept a narrowing projection that preserves the sort key");
665        assert!(
666            swapped_spm
667                .downcast_ref::<SortPreservingMergeExec>()
668                .is_some(),
669            "bare SPM should rebuild as bare SPM"
670        );
671
672        let merge_sort =
673            Arc::new(MergeSortExec::new(ordering, input, Some(1))) as Arc<dyn ExecutionPlan>;
674        let merge_projection = ProjectionExec::try_new(
675            vec![
676                (physical_col("ts", schema.as_ref())?, "ts".to_string()),
677                (physical_col("tag", schema.as_ref())?, "tag".to_string()),
678            ],
679            Arc::clone(&merge_sort),
680        )?;
681        let swapped_merge_sort = merge_sort
682            .try_swapping_with_projection(&merge_projection)?
683            .expect("MergeSortExec should accept the same projection swap as SPM");
684
685        assert!(
686            swapped_merge_sort.downcast_ref::<MergeSortExec>().is_some(),
687            "MergeSortExec must rewrap projection swaps as MergeSortExec"
688        );
689        assert!(
690            swapped_merge_sort
691                .downcast_ref::<SortPreservingMergeExec>()
692                .is_none(),
693            "MergeSortExec must not expose a bare SPM after projection swap"
694        );
695        assert_eq!(swapped_merge_sort.fetch(), Some(1));
696        assert!(
697            swapped_merge_sort.children()[0]
698                .downcast_ref::<ProjectionExec>()
699                .is_some(),
700            "the projection should move below MergeSortExec"
701        );
702        let swapped_schema = swapped_merge_sort.schema();
703        assert_eq!(
704            swapped_schema
705                .fields()
706                .iter()
707                .map(|field| field.name().as_str())
708                .collect::<Vec<_>>(),
709            vec!["ts", "tag"],
710            "swapped MergeSortExec should expose the projected schema"
711        );
712
713        let projected_ordering = LexOrdering::new([PhysicalSortExpr::new(
714            physical_col("ts", swapped_merge_sort.children()[0].schema().as_ref())?,
715            SortOptions {
716                descending: true,
717                nulls_first: false,
718            },
719        )])
720        .unwrap();
721        assert_eq!(
722            swapped_merge_sort.required_input_ordering(),
723            vec![Some(OrderingRequirements::from(projected_ordering))],
724            "projection swap must rewrite the ordering to the child projection's schema"
725        );
726        assert_eq!(
727            swapped_merge_sort.required_input_ordering(),
728            swapped_spm.required_input_ordering(),
729            "MergeSortExec projection swap should mirror SPM's ordering rewrite"
730        );
731
732        let spm_projection_without_sort_key = ProjectionExec::try_new(
733            vec![(physical_col("tag", schema.as_ref())?, "tag".to_string())],
734            Arc::clone(&bare_spm),
735        )?;
736        let merge_projection_without_sort_key = ProjectionExec::try_new(
737            vec![(physical_col("tag", schema.as_ref())?, "tag".to_string())],
738            Arc::clone(&merge_sort),
739        )?;
740        assert!(
741            bare_spm
742                .try_swapping_with_projection(&spm_projection_without_sort_key)?
743                .is_none(),
744            "SPM must reject projection swaps that drop the sort key"
745        );
746        assert!(
747            merge_sort
748                .try_swapping_with_projection(&merge_projection_without_sort_key)?
749                .is_none(),
750            "MergeSortExec should reject the same projection swap as SPM"
751        );
752
753        Ok(())
754    }
755
756    #[test]
757    fn enforce_sorting_rewrite_keeps_merge_sort_exec_opaque() {
758        let schema = Arc::new(Schema::new(vec![Field::new("ts", DataType::Int64, false)]));
759        let input = Arc::new(EmptyExec::new(schema.clone()).with_partitions(2)) as _;
760        let ordering = test_ordering(schema.as_ref());
761
762        let bare_spm = Arc::new(
763            SortPreservingMergeExec::new(ordering.clone(), Arc::clone(&input)).with_fetch(Some(1)),
764        ) as Arc<dyn ExecutionPlan>;
765        let optimized_spm = plan_with_order_breaking_variants(OrderPreservationContext::new(
766            bare_spm,
767            false,
768            vec![OrderPreservationContext::new(
769                Arc::clone(&input),
770                false,
771                vec![],
772            )],
773        ))
774        .unwrap()
775        .plan;
776        assert!(
777            optimized_spm
778                .downcast_ref::<CoalescePartitionsExec>()
779                .is_some(),
780            "this regression test must exercise EnforceSorting's bare SPM -> CoalescePartitionsExec rewrite"
781        );
782
783        let merge_sort =
784            Arc::new(MergeSortExec::new(ordering, input, Some(1))) as Arc<dyn ExecutionPlan>;
785        let optimized_merge_sort =
786            plan_with_order_breaking_variants(OrderPreservationContext::new(
787                Arc::clone(&merge_sort),
788                false,
789                vec![OrderPreservationContext::new(
790                    Arc::clone(merge_sort.children()[0]),
791                    false,
792                    vec![],
793                )],
794            ))
795            .unwrap()
796            .plan;
797        assert!(
798            optimized_merge_sort
799                .downcast_ref::<MergeSortExec>()
800                .is_some(),
801            "MergeSortExec must stay opaque to the bare SPM rewrite"
802        );
803        assert!(
804            optimized_merge_sort
805                .downcast_ref::<CoalescePartitionsExec>()
806                .is_none(),
807            "MergeSortExec(fetch) is the required distributed TopK merge stage, not an unordered coalesce"
808        );
809        assert_eq!(optimized_merge_sort.fetch(), Some(1));
810    }
811}