Skip to main content

query/
part_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//! Module for sorting input data within each [`PartitionRange`].
16//!
17//! This module defines the [`PartSortExec`] execution plan, which sorts each
18//! partition ([`PartitionRange`]) independently based on the provided physical
19//! sort expressions.
20
21use std::pin::Pin;
22use std::sync::Arc;
23use std::task::{Context, Poll};
24
25use arrow::array::{Array, ArrayRef};
26use arrow::compute::{concat, concat_batches, take_record_batch};
27use arrow_schema::{DataType, SchemaRef, TimeUnit};
28use common_recordbatch::{DfRecordBatch, DfSendableRecordBatchStream};
29use common_time::Timestamp;
30use datafusion::common::arrow::compute::sort_to_indices;
31use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation};
32use datafusion::execution::{RecordBatchStream, TaskContext};
33use datafusion::physical_plan::execution_plan::CardinalityEffect;
34use datafusion::physical_plan::filter_pushdown::{
35    ChildFilterDescription, FilterDescription, FilterPushdownPhase,
36};
37use datafusion::physical_plan::metrics::{BaselineMetrics, ExecutionPlanMetricsSet, MetricsSet};
38use datafusion::physical_plan::{
39    DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, PlanProperties,
40    apply_expression_roots,
41};
42use datafusion_common::tree_node::TreeNodeRecursion;
43use datafusion_common::{DataFusionError, ScalarValue, internal_err};
44use datafusion_expr::Operator;
45use datafusion_physical_expr::expressions::{
46    BinaryExpr, DynamicFilterPhysicalExpr, is_not_null, is_null, lit,
47};
48use datafusion_physical_expr::{PhysicalExpr, PhysicalSortExpr};
49use futures::Stream;
50use itertools::Itertools;
51use snafu::location;
52use store_api::region_engine::PartitionRange;
53
54use crate::error::Result;
55use crate::window_sort::{check_partition_range_monotonicity, project_partition_range_for_sort};
56use crate::{array_iter_helper, downcast_ts_array};
57
58/// Get the primary end of a `PartitionRange` based on sort direction.
59///
60/// - Descending: primary end is `end` (we process highest values first)
61/// - Ascending: primary end is `start` (we process lowest values first)
62fn get_primary_end(range: &PartitionRange, descending: bool) -> Timestamp {
63    if descending { range.end } else { range.start }
64}
65
66/// Group consecutive ranges by their primary end value.
67///
68/// Returns a vector of (primary_end, start_idx_inclusive, end_idx_exclusive) tuples.
69/// Ranges with the same primary end MUST be processed together because they may
70/// overlap and contain values that belong to the same "top-k" result.
71fn group_ranges_by_primary_end(
72    ranges: &[PartitionRange],
73    descending: bool,
74) -> Vec<(Timestamp, usize, usize)> {
75    if ranges.is_empty() {
76        return vec![];
77    }
78
79    let mut groups = Vec::new();
80    let mut group_start = 0;
81    let mut current_primary_end = get_primary_end(&ranges[0], descending);
82
83    for (idx, range) in ranges.iter().enumerate().skip(1) {
84        let primary_end = get_primary_end(range, descending);
85        if primary_end != current_primary_end {
86            // End current group
87            groups.push((current_primary_end, group_start, idx));
88            // Start new group
89            group_start = idx;
90            current_primary_end = primary_end;
91        }
92    }
93    // Push the last group
94    groups.push((current_primary_end, group_start, ranges.len()));
95
96    groups
97}
98
99/// Sort input within given PartitionRange
100///
101/// Partition range transitions are detected by comparing sort column values against
102/// the current range boundaries (via [`PartSortStream::try_find_next_range`]).
103/// Empty RecordBatches from upstream are tolerated but do not serve as range delimiters.
104///
105/// This operator sorts each partition independently.
106#[derive(Debug, Clone)]
107pub struct PartSortExec {
108    /// Physical sort expressions(that is, sort by timestamp)
109    expression: PhysicalSortExpr,
110    limit: Option<usize>,
111    input: Arc<dyn ExecutionPlan>,
112    /// Execution metrics
113    metrics: ExecutionPlanMetricsSet,
114    partition_ranges: Vec<Vec<PartitionRange>>,
115    properties: Arc<PlanProperties>,
116    dynamic_filter: Option<Arc<DynamicFilterPhysicalExpr>>,
117}
118
119impl PartSortExec {
120    pub fn try_new(
121        expression: PhysicalSortExpr,
122        limit: Option<usize>,
123        partition_ranges: Vec<Vec<PartitionRange>>,
124        input: Arc<dyn ExecutionPlan>,
125    ) -> Result<Self> {
126        check_partition_range_monotonicity(&partition_ranges, expression.options.descending)?;
127
128        let metrics = ExecutionPlanMetricsSet::new();
129        let properties = input.properties();
130        let properties = Arc::new(PlanProperties::new(
131            input.equivalence_properties().clone(),
132            input.output_partitioning().clone(),
133            properties.emission_type,
134            properties.boundedness,
135        ));
136
137        let dynamic_filter = Self::new_dynamic_filter(&expression, limit);
138
139        Ok(Self {
140            expression,
141            limit,
142            input,
143            metrics,
144            partition_ranges,
145            properties,
146            dynamic_filter,
147        })
148    }
149
150    fn new_dynamic_filter(
151        expression: &PhysicalSortExpr,
152        limit: Option<usize>,
153    ) -> Option<Arc<DynamicFilterPhysicalExpr>> {
154        limit.map(|_| {
155            Arc::new(DynamicFilterPhysicalExpr::new(
156                vec![expression.expr.clone()],
157                lit(true),
158            ))
159        })
160    }
161
162    pub fn to_stream(
163        &self,
164        context: Arc<TaskContext>,
165        partition: usize,
166    ) -> datafusion_common::Result<DfSendableRecordBatchStream> {
167        let input_stream: DfSendableRecordBatchStream =
168            self.input.execute(partition, context.clone())?;
169
170        if partition >= self.partition_ranges.len() {
171            internal_err!(
172                "Partition index out of range: {} >= {} at {}",
173                partition,
174                self.partition_ranges.len(),
175                snafu::location!()
176            )?;
177        }
178
179        let df_stream = Box::pin(PartSortStream::new(
180            context,
181            self,
182            self.limit,
183            input_stream,
184            self.partition_ranges[partition].clone(),
185            partition,
186        )?) as _;
187
188        Ok(df_stream)
189    }
190}
191
192impl DisplayAs for PartSortExec {
193    fn fmt_as(&self, _t: DisplayFormatType, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
194        write!(
195            f,
196            "PartSortExec: expr={} num_ranges={}",
197            self.expression,
198            self.partition_ranges.len(),
199        )?;
200        if let Some(limit) = self.limit {
201            write!(f, " limit={}", limit)?;
202        }
203        Ok(())
204    }
205}
206
207impl ExecutionPlan for PartSortExec {
208    fn name(&self) -> &str {
209        "PartSortExec"
210    }
211
212    fn schema(&self) -> SchemaRef {
213        self.input.schema()
214    }
215
216    fn properties(&self) -> &Arc<PlanProperties> {
217        &self.properties
218    }
219
220    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
221        vec![&self.input]
222    }
223
224    fn apply_expressions(
225        &self,
226        f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> datafusion_common::Result<TreeNodeRecursion>,
227    ) -> datafusion_common::Result<TreeNodeRecursion> {
228        let dynamic_filter = self
229            .dynamic_filter
230            .as_ref()
231            .map(|filter| filter.clone() as Arc<dyn PhysicalExpr>);
232        apply_expression_roots(
233            std::iter::once(&self.expression.expr).chain(dynamic_filter.as_ref()),
234            f,
235        )
236    }
237
238    fn dynamic_expressions_produced(&self) -> Vec<Arc<dyn PhysicalExpr>> {
239        self.dynamic_filter
240            .iter()
241            .map(|filter| filter.clone() as Arc<dyn PhysicalExpr>)
242            .collect()
243    }
244
245    fn with_new_children(
246        self: Arc<Self>,
247        children: Vec<Arc<dyn ExecutionPlan>>,
248    ) -> datafusion_common::Result<Arc<dyn ExecutionPlan>> {
249        let new_input = if let Some(first) = children.first() {
250            first
251        } else {
252            internal_err!("No children found")?
253        };
254        let mut new_exec = self.as_ref().clone();
255        new_exec.input = new_input.clone();
256        new_exec.properties = new_input.properties().clone();
257        Ok(Arc::new(new_exec))
258    }
259
260    fn execute(
261        &self,
262        partition: usize,
263        context: Arc<TaskContext>,
264    ) -> datafusion_common::Result<DfSendableRecordBatchStream> {
265        self.to_stream(context, partition)
266    }
267
268    fn metrics(&self) -> Option<MetricsSet> {
269        Some(self.metrics.clone_inner())
270    }
271
272    /// # Explain
273    ///
274    /// This plan needs to be executed on each partition independently,
275    /// and is expected to run directly on storage engine's output
276    /// distribution / partition.
277    fn benefits_from_input_partitioning(&self) -> Vec<bool> {
278        vec![false]
279    }
280
281    fn cardinality_effect(&self) -> CardinalityEffect {
282        if self.limit.is_none() {
283            CardinalityEffect::Equal
284        } else {
285            CardinalityEffect::LowerEqual
286        }
287    }
288
289    fn gather_filters_for_pushdown(
290        &self,
291        phase: FilterPushdownPhase,
292        parent_filters: Vec<Arc<dyn PhysicalExpr>>,
293        config: &datafusion::config::ConfigOptions,
294    ) -> datafusion_common::Result<FilterDescription> {
295        if !matches!(phase, FilterPushdownPhase::Post) {
296            return FilterDescription::from_children(parent_filters, &self.children());
297        }
298
299        let mut child = ChildFilterDescription::from_child(&parent_filters, &self.input)?;
300        if let Some(filter) = &self.dynamic_filter
301            && config.optimizer.enable_topk_dynamic_filter_pushdown
302        {
303            let filter: Arc<dyn PhysicalExpr> = filter.clone();
304            child = child.with_self_filter(filter);
305        }
306        Ok(FilterDescription::new().with_child(child))
307    }
308
309    fn reset_state(self: Arc<Self>) -> datafusion_common::Result<Arc<dyn ExecutionPlan>> {
310        let dynamic_filter = Self::new_dynamic_filter(&self.expression, self.limit);
311        Ok(Arc::new(Self {
312            expression: self.expression.clone(),
313            limit: self.limit,
314            input: self.input.clone(),
315            metrics: self.metrics.clone(),
316            partition_ranges: self.partition_ranges.clone(),
317            properties: self.properties.clone(),
318            dynamic_filter,
319        }))
320    }
321}
322
323enum PartSortBuffer {
324    All(Vec<DfRecordBatch>),
325    TopK(Vec<DfRecordBatch>),
326}
327
328#[derive(Clone, Debug, PartialEq, Eq)]
329enum TopKThreshold {
330    Null,
331    Value(i64),
332}
333
334impl PartSortBuffer {
335    pub fn is_empty(&self) -> bool {
336        match self {
337            PartSortBuffer::All(v) => v.is_empty(),
338            PartSortBuffer::TopK(v) => v.is_empty(),
339        }
340    }
341
342    pub fn num_rows(&self) -> usize {
343        match self {
344            PartSortBuffer::All(v) => v.iter().map(|batch| batch.num_rows()).sum(),
345            PartSortBuffer::TopK(v) => v.iter().map(|batch| batch.num_rows()).sum(),
346        }
347    }
348}
349
350struct PartSortStream {
351    /// Memory pool for this stream
352    reservation: MemoryReservation,
353    buffer: PartSortBuffer,
354    expression: PhysicalSortExpr,
355    limit: Option<usize>,
356    input: DfSendableRecordBatchStream,
357    input_complete: bool,
358    schema: SchemaRef,
359    partition_ranges: Vec<PartitionRange>,
360    #[allow(dead_code)] // this is used under #[debug_assertions]
361    partition: usize,
362    cur_part_idx: usize,
363    evaluating_batch: Option<DfRecordBatch>,
364    metrics: BaselineMetrics,
365    dynamic_filter: Option<Arc<DynamicFilterPhysicalExpr>>,
366    dynamic_filter_threshold: Option<TopKThreshold>,
367    /// Groups of ranges by primary end: (primary_end, start_idx_inclusive, end_idx_exclusive).
368    /// Ranges in the same group must be processed together before outputting results.
369    range_groups: Vec<(Timestamp, usize, usize)>,
370    /// Current group being processed (index into range_groups).
371    cur_group_idx: usize,
372}
373
374impl PartSortStream {
375    fn new(
376        context: Arc<TaskContext>,
377        sort: &PartSortExec,
378        limit: Option<usize>,
379        input: DfSendableRecordBatchStream,
380        partition_ranges: Vec<PartitionRange>,
381        partition: usize,
382    ) -> datafusion_common::Result<Self> {
383        let buffer = if limit.is_some() {
384            PartSortBuffer::TopK(Vec::new())
385        } else {
386            PartSortBuffer::All(Vec::new())
387        };
388
389        // Compute range groups by primary end
390        let descending = sort.expression.options.descending;
391        let range_groups = group_ranges_by_primary_end(&partition_ranges, descending);
392
393        Ok(Self {
394            reservation: MemoryConsumer::new("PartSortStream".to_string())
395                .register(&context.runtime_env().memory_pool),
396            buffer,
397            expression: sort.expression.clone(),
398            limit,
399            input,
400            input_complete: false,
401            schema: sort.input.schema(),
402            partition_ranges,
403            partition,
404            cur_part_idx: 0,
405            evaluating_batch: None,
406            metrics: BaselineMetrics::new(&sort.metrics, partition),
407            dynamic_filter: sort.dynamic_filter.clone(),
408            dynamic_filter_threshold: None,
409            range_groups,
410            cur_group_idx: 0,
411        })
412    }
413}
414
415macro_rules! array_check_helper {
416    ($t:ty, $unit:expr, $arr:expr, $cur_range:expr, $min_max_idx:expr) => {{
417            if $cur_range.start.unit().as_arrow_time_unit() != $unit
418            || $cur_range.end.unit().as_arrow_time_unit() != $unit
419        {
420            internal_err!(
421                "PartitionRange unit mismatch, expect {:?}, found {:?}",
422                $cur_range.start.unit(),
423                $unit
424            )?;
425        }
426        let arr = $arr
427            .as_any()
428            .downcast_ref::<arrow::array::PrimitiveArray<$t>>()
429            .unwrap();
430
431        let min = arr.value($min_max_idx.0);
432        let max = arr.value($min_max_idx.1);
433        let (min, max) = if min < max{
434            (min, max)
435        } else {
436            (max, min)
437        };
438        let cur_min = $cur_range.start.value();
439        let cur_max = $cur_range.end.value();
440        // note that PartitionRange is left inclusive and right exclusive
441        if !(min >= cur_min && max < cur_max) {
442            internal_err!(
443                "Sort column min/max value out of partition range: sort_column.min_max=[{:?}, {:?}] not in PartitionRange=[{:?}, {:?}]",
444                min,
445                max,
446                cur_min,
447                cur_max
448            )?;
449        }
450    }};
451}
452
453macro_rules! threshold_helper {
454    ($t:ty, $unit:expr, $arr:expr, $threshold_idx:expr) => {{
455        let arr = $arr
456            .as_any()
457            .downcast_ref::<arrow::array::PrimitiveArray<$t>>()
458            .unwrap();
459        if arr.is_null($threshold_idx) {
460            TopKThreshold::Null
461        } else {
462            TopKThreshold::Value(arr.value($threshold_idx))
463        }
464    }};
465}
466
467impl PartSortStream {
468    /// check whether the sort column's min/max value is within the current group's effective range.
469    /// For group-based processing, data from multiple ranges with the same primary end
470    /// is accumulated together, so we check against the union of all ranges in the group.
471    fn check_in_range(
472        &self,
473        sort_column: &ArrayRef,
474        min_max_idx: (usize, usize),
475    ) -> datafusion_common::Result<()> {
476        // Use the group's effective range instead of the current partition range
477        let Some(cur_range) = self.get_current_group_effective_range() else {
478            internal_err!(
479                "No effective range for current group {} at {}",
480                self.cur_group_idx,
481                snafu::location!()
482            )?
483        };
484        let cur_range = project_partition_range_for_sort(cur_range, sort_column.data_type())?;
485
486        downcast_ts_array!(
487            sort_column.data_type() => (array_check_helper, sort_column, cur_range, min_max_idx),
488            _ => internal_err!(
489                "Unsupported data type for sort column: {:?}",
490                sort_column.data_type()
491            )?,
492        );
493
494        Ok(())
495    }
496
497    /// Try find data whose value exceeds the current partition range.
498    ///
499    /// Returns `None` if no such data is found, and `Some(idx)` where idx points to
500    /// the first data that exceeds the current partition range.
501    fn try_find_next_range(
502        &self,
503        sort_column: &ArrayRef,
504    ) -> datafusion_common::Result<Option<usize>> {
505        if sort_column.is_empty() {
506            return Ok(None);
507        }
508
509        // check if the current partition index is out of range
510        if self.cur_part_idx >= self.partition_ranges.len() {
511            internal_err!(
512                "Partition index out of range: {} >= {} at {}",
513                self.cur_part_idx,
514                self.partition_ranges.len(),
515                snafu::location!()
516            )?;
517        }
518        let cur_range = project_partition_range_for_sort(
519            self.partition_ranges[self.cur_part_idx],
520            sort_column.data_type(),
521        )?;
522
523        let sort_column_iter = downcast_ts_array!(
524            sort_column.data_type() => (array_iter_helper, sort_column),
525            _ => internal_err!(
526                "Unsupported data type for sort column: {:?}",
527                sort_column.data_type()
528            )?,
529        );
530
531        for (idx, val) in sort_column_iter {
532            // ignore vacant time index data
533            if let Some(val) = val
534                && (val >= cur_range.end.value() || val < cur_range.start.value())
535            {
536                return Ok(Some(idx));
537            }
538        }
539
540        Ok(None)
541    }
542
543    fn push_buffer(
544        &mut self,
545        batch: DfRecordBatch,
546        sort_data_type: &DataType,
547    ) -> datafusion_common::Result<()> {
548        let topk = matches!(self.buffer, PartSortBuffer::TopK(_));
549        match &mut self.buffer {
550            PartSortBuffer::All(v) => v.push(batch),
551            PartSortBuffer::TopK(v) => v.push(batch),
552        }
553
554        if topk {
555            let threshold = self.compact_topk_buffer(sort_data_type)?;
556            self.update_dynamic_filter(sort_data_type, threshold)?;
557        }
558
559        Ok(())
560    }
561
562    fn compact_topk_buffer(
563        &mut self,
564        sort_data_type: &DataType,
565    ) -> datafusion_common::Result<Option<TopKThreshold>> {
566        let Some(limit) = self.limit else {
567            return Ok(None);
568        };
569
570        let PartSortBuffer::TopK(buffer) =
571            std::mem::replace(&mut self.buffer, PartSortBuffer::TopK(Vec::new()))
572        else {
573            return Ok(None);
574        };
575
576        if limit == 0 || buffer.is_empty() {
577            self.buffer = PartSortBuffer::TopK(Vec::new());
578            return Ok(None);
579        }
580
581        let total_rows: usize = buffer.iter().map(|batch| batch.num_rows()).sum();
582        if total_rows <= limit {
583            self.buffer = PartSortBuffer::TopK(buffer);
584            return Ok(None);
585        }
586
587        let topk = self.sort_record_batches(&buffer, Some(limit), false)?;
588        let threshold = self.threshold_from_sorted_batch(&topk, sort_data_type)?;
589        self.buffer = if topk.num_rows() == 0 {
590            PartSortBuffer::TopK(Vec::new())
591        } else {
592            PartSortBuffer::TopK(vec![topk])
593        };
594
595        Ok(threshold)
596    }
597
598    fn threshold_from_sorted_batch(
599        &self,
600        batch: &DfRecordBatch,
601        sort_data_type: &DataType,
602    ) -> datafusion_common::Result<Option<TopKThreshold>> {
603        if batch.num_rows() == 0 {
604            return Ok(None);
605        }
606
607        let threshold_idx = batch.num_rows() - 1;
608        let sort_column = self.expression.evaluate_to_sort_column(batch)?.values;
609        let threshold = downcast_ts_array!(
610            sort_data_type => (threshold_helper, sort_column, threshold_idx),
611            _ => internal_err!(
612                "Unsupported data type for sort column: {:?}",
613                sort_data_type
614            )?,
615        );
616
617        Ok(Some(threshold))
618    }
619
620    fn topk_threshold(
621        &self,
622        sort_data_type: &arrow_schema::DataType,
623    ) -> datafusion_common::Result<Option<TopKThreshold>> {
624        let Some(limit) = self.limit else {
625            return Ok(None);
626        };
627
628        if limit == 0 || self.buffer.num_rows() < limit {
629            return Ok(None);
630        }
631
632        let buffer = match &self.buffer {
633            PartSortBuffer::All(buffer) | PartSortBuffer::TopK(buffer) => buffer,
634        };
635        let mut sort_columns = Vec::with_capacity(buffer.len());
636        let mut opt = None;
637        for batch in buffer {
638            let sort_column = self.expression.evaluate_to_sort_column(batch)?;
639            opt = opt.or(sort_column.options);
640            sort_columns.push(sort_column.values);
641        }
642
643        let sort_column =
644            concat(&sort_columns.iter().map(|a| a.as_ref()).collect_vec()).map_err(|e| {
645                DataFusionError::ArrowError(
646                    Box::new(e),
647                    Some(format!("Fail to concat sort columns at {}", location!())),
648                )
649            })?;
650
651        let indices = sort_to_indices(&sort_column, opt, Some(limit)).map_err(|e| {
652            DataFusionError::ArrowError(
653                Box::new(e),
654                Some(format!("Fail to sort to indices at {}", location!())),
655            )
656        })?;
657
658        if indices.len() < limit {
659            return Ok(None);
660        }
661
662        let threshold_idx = indices.value(indices.len() - 1) as usize;
663        let threshold = downcast_ts_array!(
664            sort_data_type => (threshold_helper, sort_column, threshold_idx),
665            _ => internal_err!(
666                "Unsupported data type for sort column: {:?}",
667                sort_data_type
668            )?,
669        );
670
671        Ok(Some(threshold))
672    }
673
674    fn threshold_scalar_value(
675        sort_data_type: &DataType,
676        threshold: &TopKThreshold,
677    ) -> datafusion_common::Result<ScalarValue> {
678        let value = match threshold {
679            TopKThreshold::Null => None,
680            TopKThreshold::Value(value) => Some(*value),
681        };
682
683        let scalar = match sort_data_type {
684            DataType::Timestamp(TimeUnit::Second, tz) => {
685                ScalarValue::TimestampSecond(value, tz.clone())
686            }
687            DataType::Timestamp(TimeUnit::Millisecond, tz) => {
688                ScalarValue::TimestampMillisecond(value, tz.clone())
689            }
690            DataType::Timestamp(TimeUnit::Microsecond, tz) => {
691                ScalarValue::TimestampMicrosecond(value, tz.clone())
692            }
693            DataType::Timestamp(TimeUnit::Nanosecond, tz) => {
694                ScalarValue::TimestampNanosecond(value, tz.clone())
695            }
696            _ => internal_err!(
697                "Unsupported data type for sort column: {:?}",
698                sort_data_type
699            )?,
700        };
701
702        Ok(scalar)
703    }
704
705    fn build_dynamic_filter_expr(
706        &self,
707        sort_data_type: &DataType,
708        threshold: &TopKThreshold,
709    ) -> datafusion_common::Result<Arc<dyn PhysicalExpr>> {
710        let op = if self.expression.options.descending {
711            Operator::Gt
712        } else {
713            Operator::Lt
714        };
715        let value_null = matches!(threshold, TopKThreshold::Null);
716        let value = Self::threshold_scalar_value(sort_data_type, threshold)?;
717        let comparison: Arc<dyn PhysicalExpr> = Arc::new(BinaryExpr::new(
718            self.expression.expr.clone(),
719            op,
720            lit(value),
721        ));
722
723        match (self.expression.options.nulls_first, value_null) {
724            (true, true) => Ok(lit(false)),
725            (true, false) => Ok(Arc::new(BinaryExpr::new(
726                is_null(self.expression.expr.clone())?,
727                Operator::Or,
728                comparison,
729            ))),
730            (false, true) => is_not_null(self.expression.expr.clone()),
731            (false, false) => Ok(comparison),
732        }
733    }
734
735    fn update_dynamic_filter(
736        &mut self,
737        sort_data_type: &DataType,
738        threshold: Option<TopKThreshold>,
739    ) -> datafusion_common::Result<()> {
740        let Some(filter) = &self.dynamic_filter else {
741            return Ok(());
742        };
743
744        let threshold = if let Some(threshold) = threshold {
745            threshold
746        } else {
747            let Some(threshold) = self.topk_threshold(sort_data_type)? else {
748                return Ok(());
749            };
750            threshold
751        };
752
753        if self.dynamic_filter_threshold.as_ref() == Some(&threshold) {
754            return Ok(());
755        }
756
757        let predicate = self.build_dynamic_filter_expr(sort_data_type, &threshold)?;
758        filter.update(predicate)?;
759        self.dynamic_filter_threshold = Some(threshold);
760
761        Ok(())
762    }
763
764    /// Returns true when all rows in the next group are guaranteed to be worse
765    /// than the current top-k threshold.
766    fn can_stop_before_group(
767        &self,
768        group_idx: usize,
769        sort_data_type: &arrow_schema::DataType,
770    ) -> datafusion_common::Result<bool> {
771        if group_idx >= self.range_groups.len() {
772            return Ok(false);
773        }
774
775        let threshold = if let Some(threshold) = &self.dynamic_filter_threshold {
776            threshold.clone()
777        } else {
778            let Some(threshold) = self.topk_threshold(sort_data_type)? else {
779                return Ok(false);
780            };
781            threshold
782        };
783
784        let (_, start_idx, _) = self.range_groups[group_idx];
785        let next_range =
786            project_partition_range_for_sort(self.partition_ranges[start_idx], sort_data_type)?;
787        let descending = self.expression.options.descending;
788        let next_primary = get_primary_end(&next_range, descending).value();
789
790        let can_stop = match threshold {
791            // When the k-th element is NULL:
792            // - nulls_first=true: nulls are the best values (Arrow sorts NULLs first
793            //   regardless of ASC/DESC), so top-k is already optimal → stop early.
794            // - nulls_first=false: nulls are the worst, non-null values from the
795            //   next group could displace them → continue reading.
796            TopKThreshold::Null => self.expression.options.nulls_first,
797            TopKThreshold::Value(value) => {
798                if descending {
799                    value >= next_primary
800                } else {
801                    value < next_primary
802                }
803            }
804        };
805
806        Ok(can_stop)
807    }
808
809    /// Check if the given partition index is within the current group.
810    fn is_in_current_group(&self, part_idx: usize) -> bool {
811        if self.cur_group_idx >= self.range_groups.len() {
812            return false;
813        }
814        let (_, start, end) = self.range_groups[self.cur_group_idx];
815        part_idx >= start && part_idx < end
816    }
817
818    /// Advance to the next group. Returns true if there is a next group.
819    fn advance_to_next_group(&mut self) -> bool {
820        self.cur_group_idx += 1;
821        self.cur_group_idx < self.range_groups.len()
822    }
823
824    /// Get the effective range for the current group.
825    /// For a group of ranges with the same primary end, the effective range is
826    /// the union of all ranges in the group.
827    fn get_current_group_effective_range(&self) -> Option<PartitionRange> {
828        if self.cur_group_idx >= self.range_groups.len() {
829            return None;
830        }
831        let (_, start_idx, end_idx) = self.range_groups[self.cur_group_idx];
832        if start_idx >= end_idx || start_idx >= self.partition_ranges.len() {
833            return None;
834        }
835
836        let ranges_in_group =
837            &self.partition_ranges[start_idx..end_idx.min(self.partition_ranges.len())];
838        if ranges_in_group.is_empty() {
839            return None;
840        }
841
842        // Compute union of all ranges in the group
843        let mut min_start = ranges_in_group[0].start;
844        let mut max_end = ranges_in_group[0].end;
845        for range in ranges_in_group.iter().skip(1) {
846            if range.start < min_start {
847                min_start = range.start;
848            }
849            if range.end > max_end {
850                max_end = range.end;
851            }
852        }
853
854        Some(PartitionRange {
855            start: min_start,
856            end: max_end,
857            num_rows: 0,   // Not used for validation
858            identifier: 0, // Not used for validation
859        })
860    }
861
862    /// Sort and clear the buffer and return the sorted record batch
863    ///
864    /// this function will return a empty record batch if the buffer is empty
865    fn sort_buffer(&mut self) -> datafusion_common::Result<DfRecordBatch> {
866        match &mut self.buffer {
867            PartSortBuffer::All(_) => self.sort_all_buffer(),
868            PartSortBuffer::TopK(_) => self.sort_topk_buffer(),
869        }
870    }
871
872    fn sort_record_batches(
873        &mut self,
874        buffer: &[DfRecordBatch],
875        limit: Option<usize>,
876        check_range: bool,
877    ) -> datafusion_common::Result<DfRecordBatch> {
878        if buffer.is_empty() {
879            return Ok(DfRecordBatch::new_empty(self.schema.clone()));
880        }
881
882        let mut sort_columns = Vec::with_capacity(buffer.len());
883        let mut opt = None;
884        for batch in buffer.iter() {
885            let sort_column = self.expression.evaluate_to_sort_column(batch)?;
886            opt = opt.or(sort_column.options);
887            sort_columns.push(sort_column.values);
888        }
889
890        let sort_column =
891            concat(&sort_columns.iter().map(|a| a.as_ref()).collect_vec()).map_err(|e| {
892                DataFusionError::ArrowError(
893                    Box::new(e),
894                    Some(format!("Fail to concat sort columns at {}", location!())),
895                )
896            })?;
897
898        let indices = sort_to_indices(&sort_column, opt, limit).map_err(|e| {
899            DataFusionError::ArrowError(
900                Box::new(e),
901                Some(format!("Fail to sort to indices at {}", location!())),
902            )
903        })?;
904        if indices.is_empty() {
905            return Ok(DfRecordBatch::new_empty(self.schema.clone()));
906        }
907
908        if check_range {
909            self.check_in_range(
910                &sort_column,
911                (
912                    indices.value(0) as usize,
913                    indices.value(indices.len() - 1) as usize,
914                ),
915            )
916            .inspect_err(|_e| {
917                #[cfg(debug_assertions)]
918                common_telemetry::error!(
919                    "Fail to check sort column in range at {}, current_idx: {}, num_rows: {}, err: {}",
920                    self.partition,
921                    self.cur_part_idx,
922                    sort_column.len(),
923                    _e
924                );
925            })?;
926        }
927
928        // reserve memory for the concat input and sorted output
929        let total_mem: usize = buffer.iter().map(|r| r.get_array_memory_size()).sum();
930        self.reservation.try_grow(total_mem * 2)?;
931
932        let full_input = concat_batches(&self.schema, buffer).map_err(|e| {
933            DataFusionError::ArrowError(
934                Box::new(e),
935                Some(format!(
936                    "Fail to concat input batches when sorting at {}",
937                    location!()
938                )),
939            )
940        })?;
941
942        let sorted = take_record_batch(&full_input, &indices).map_err(|e| {
943            DataFusionError::ArrowError(
944                Box::new(e),
945                Some(format!(
946                    "Fail to take result record batch when sorting at {}",
947                    location!()
948                )),
949            )
950        })?;
951
952        drop(full_input);
953        // here remove both buffer and full_input memory
954        self.reservation.shrink(2 * total_mem);
955        Ok(sorted)
956    }
957
958    /// Internal method for sorting `All` buffer (without limit).
959    fn sort_all_buffer(&mut self) -> datafusion_common::Result<DfRecordBatch> {
960        let PartSortBuffer::All(buffer) =
961            std::mem::replace(&mut self.buffer, PartSortBuffer::All(Vec::new()))
962        else {
963            unreachable!()
964        };
965
966        self.sort_record_batches(&buffer, self.limit, self.limit.is_none())
967    }
968
969    fn sort_topk_buffer(&mut self) -> datafusion_common::Result<DfRecordBatch> {
970        let PartSortBuffer::TopK(buffer) =
971            std::mem::replace(&mut self.buffer, PartSortBuffer::TopK(Vec::new()))
972        else {
973            unreachable!()
974        };
975
976        self.sort_record_batches(&buffer, self.limit, false)
977    }
978
979    /// Sorts current buffer and returns `None` when there is nothing to emit.
980    fn sorted_buffer_if_non_empty(&mut self) -> datafusion_common::Result<Option<DfRecordBatch>> {
981        if self.buffer.is_empty() {
982            return Ok(None);
983        }
984
985        let sorted = self.sort_buffer()?;
986        if sorted.num_rows() == 0 {
987            Ok(None)
988        } else {
989            Ok(Some(sorted))
990        }
991    }
992
993    fn mark_dynamic_filter_complete(&self) {
994        if let Some(filter) = &self.dynamic_filter {
995            filter.mark_complete();
996        }
997    }
998
999    /// Try to split the input batch if it contains data that exceeds the current partition range.
1000    ///
1001    /// When the input batch contains data that exceeds the current partition range, this function
1002    /// will split the input batch into two parts, the first part is within the current partition
1003    /// range will be merged and sorted with previous buffer, and the second part will be registered
1004    /// to `evaluating_batch` for next polling.
1005    ///
1006    /// **Group-based processing**: Ranges with the same primary end are grouped together.
1007    /// We only sort and output when transitioning to a NEW group, not when moving between
1008    /// ranges within the same group.
1009    ///
1010    /// Returns `None` if the input batch is empty or fully within the current partition range
1011    /// (or we're still collecting data within the same group), and `Some(batch)` when we've
1012    /// completed a group and have sorted output. When operating in limit mode, this
1013    /// function will not emit intermediate batches; it only prepares state for a single final
1014    /// output.
1015    fn split_batch(
1016        &mut self,
1017        batch: DfRecordBatch,
1018    ) -> datafusion_common::Result<Option<DfRecordBatch>> {
1019        if self.limit.is_some() {
1020            self.split_batch_topk(batch)?;
1021            return Ok(None);
1022        }
1023
1024        self.split_batch_all(batch)
1025    }
1026
1027    /// Specialized splitting logic for limit mode.
1028    ///
1029    /// We only emit once when input is fully consumed.
1030    /// When the buffer is fulfilled and we are about to enter a new group, we stop consuming
1031    /// further ranges.
1032    fn split_batch_topk(&mut self, batch: DfRecordBatch) -> datafusion_common::Result<()> {
1033        if batch.num_rows() == 0 {
1034            return Ok(());
1035        }
1036
1037        let sort_column = self
1038            .expression
1039            .expr
1040            .evaluate(&batch)?
1041            .into_array(batch.num_rows())?;
1042
1043        let next_range_idx = self.try_find_next_range(&sort_column)?;
1044        let Some(idx) = next_range_idx else {
1045            self.push_buffer(batch, sort_column.data_type())?;
1046            // keep polling input for next batch
1047            return Ok(());
1048        };
1049
1050        let this_range = batch.slice(0, idx);
1051        let remaining_range = batch.slice(idx, batch.num_rows() - idx);
1052        if this_range.num_rows() != 0 {
1053            self.push_buffer(this_range, sort_column.data_type())?;
1054        }
1055
1056        // Step to next proper PartitionRange
1057        self.cur_part_idx += 1;
1058
1059        // If we've processed all partitions, mark completion.
1060        if self.cur_part_idx >= self.partition_ranges.len() {
1061            debug_assert!(remaining_range.num_rows() == 0);
1062            self.input_complete = true;
1063            return Ok(());
1064        }
1065
1066        // Check if we're still in the same group
1067        let in_same_group = self.is_in_current_group(self.cur_part_idx);
1068
1069        if !in_same_group {
1070            let next_group_idx = self.cur_group_idx + 1;
1071            if self.can_stop_before_group(next_group_idx, sort_column.data_type())? {
1072                self.input_complete = true;
1073                return Ok(());
1074            }
1075            self.advance_to_next_group();
1076        }
1077
1078        let next_sort_column = sort_column.slice(idx, batch.num_rows() - idx);
1079        if self.try_find_next_range(&next_sort_column)?.is_some() {
1080            // remaining batch still contains data that exceeds the current partition range
1081            // register the remaining batch for next polling
1082            self.evaluating_batch = Some(remaining_range);
1083        } else if remaining_range.num_rows() != 0 {
1084            // remaining batch is within the current partition range
1085            // push to the buffer and continue polling
1086            self.push_buffer(remaining_range, sort_column.data_type())?;
1087        }
1088
1089        Ok(())
1090    }
1091
1092    fn split_batch_all(
1093        &mut self,
1094        batch: DfRecordBatch,
1095    ) -> datafusion_common::Result<Option<DfRecordBatch>> {
1096        if batch.num_rows() == 0 {
1097            return Ok(None);
1098        }
1099
1100        let sort_column = self
1101            .expression
1102            .expr
1103            .evaluate(&batch)?
1104            .into_array(batch.num_rows())?;
1105
1106        let next_range_idx = self.try_find_next_range(&sort_column)?;
1107        let Some(idx) = next_range_idx else {
1108            self.push_buffer(batch, sort_column.data_type())?;
1109            // keep polling input for next batch
1110            return Ok(None);
1111        };
1112
1113        let this_range = batch.slice(0, idx);
1114        let remaining_range = batch.slice(idx, batch.num_rows() - idx);
1115        if this_range.num_rows() != 0 {
1116            self.push_buffer(this_range, sort_column.data_type())?;
1117        }
1118
1119        // Step to next proper PartitionRange
1120        self.cur_part_idx += 1;
1121
1122        // If we've processed all partitions, sort and output
1123        if self.cur_part_idx >= self.partition_ranges.len() {
1124            // assert there is no data beyond the last partition range (remaining is empty).
1125            debug_assert!(remaining_range.num_rows() == 0);
1126
1127            // Sort and output the final group
1128            return self.sorted_buffer_if_non_empty();
1129        }
1130
1131        // Check if we're still in the same group
1132        if self.is_in_current_group(self.cur_part_idx) {
1133            // Same group - don't sort yet, keep collecting
1134            let next_sort_column = sort_column.slice(idx, batch.num_rows() - idx);
1135            if self.try_find_next_range(&next_sort_column)?.is_some() {
1136                // remaining batch still contains data that exceeds the current partition range
1137                self.evaluating_batch = Some(remaining_range);
1138            } else {
1139                // remaining batch is within the current partition range
1140                if remaining_range.num_rows() != 0 {
1141                    self.push_buffer(remaining_range, sort_column.data_type())?;
1142                }
1143            }
1144            // Return None to continue collecting within the same group
1145            return Ok(None);
1146        }
1147
1148        // Transitioning to a new group - sort current group and output
1149        let sorted_batch = self.sorted_buffer_if_non_empty()?;
1150        self.advance_to_next_group();
1151
1152        let next_sort_column = sort_column.slice(idx, batch.num_rows() - idx);
1153        if self.try_find_next_range(&next_sort_column)?.is_some() {
1154            // remaining batch still contains data that exceeds the current partition range
1155            // register the remaining batch for next polling
1156            self.evaluating_batch = Some(remaining_range);
1157        } else {
1158            // remaining batch is within the current partition range
1159            // push to the buffer and continue polling
1160            if remaining_range.num_rows() != 0 {
1161                self.push_buffer(remaining_range, sort_column.data_type())?;
1162            }
1163        }
1164
1165        Ok(sorted_batch)
1166    }
1167
1168    pub fn poll_next_inner(
1169        mut self: Pin<&mut Self>,
1170        cx: &mut Context<'_>,
1171    ) -> Poll<Option<datafusion_common::Result<DfRecordBatch>>> {
1172        loop {
1173            if self.input_complete {
1174                if let Some(sorted_batch) = self.sorted_buffer_if_non_empty()? {
1175                    self.mark_dynamic_filter_complete();
1176                    return Poll::Ready(Some(Ok(sorted_batch)));
1177                }
1178                self.mark_dynamic_filter_complete();
1179                return Poll::Ready(None);
1180            }
1181
1182            // if there is a remaining batch being evaluated from last run,
1183            // split on it instead of fetching new batch
1184            if let Some(evaluating_batch) = self.evaluating_batch.take()
1185                && evaluating_batch.num_rows() != 0
1186            {
1187                // Check if we've already processed all partitions
1188                if self.cur_part_idx >= self.partition_ranges.len() {
1189                    // All partitions processed, discard remaining data
1190                    if let Some(sorted_batch) = self.sorted_buffer_if_non_empty()? {
1191                        self.mark_dynamic_filter_complete();
1192                        return Poll::Ready(Some(Ok(sorted_batch)));
1193                    }
1194                    self.mark_dynamic_filter_complete();
1195                    return Poll::Ready(None);
1196                }
1197
1198                if let Some(sorted_batch) = self.split_batch(evaluating_batch)? {
1199                    return Poll::Ready(Some(Ok(sorted_batch)));
1200                }
1201                continue;
1202            }
1203
1204            // fetch next batch from input
1205            let res = self.input.as_mut().poll_next(cx);
1206            match res {
1207                Poll::Ready(Some(Ok(batch))) => {
1208                    if let Some(sorted_batch) = self.split_batch(batch)? {
1209                        return Poll::Ready(Some(Ok(sorted_batch)));
1210                    }
1211                }
1212                // input stream end, mark and continue
1213                Poll::Ready(None) => {
1214                    self.input_complete = true;
1215                }
1216                Poll::Ready(Some(Err(e))) => return Poll::Ready(Some(Err(e))),
1217                Poll::Pending => return Poll::Pending,
1218            }
1219        }
1220    }
1221}
1222
1223impl Stream for PartSortStream {
1224    type Item = datafusion_common::Result<DfRecordBatch>;
1225
1226    fn poll_next(
1227        mut self: Pin<&mut Self>,
1228        cx: &mut Context<'_>,
1229    ) -> Poll<Option<datafusion_common::Result<DfRecordBatch>>> {
1230        let result = self.as_mut().poll_next_inner(cx);
1231        self.metrics.record_poll(result)
1232    }
1233}
1234
1235impl RecordBatchStream for PartSortStream {
1236    fn schema(&self) -> SchemaRef {
1237        self.schema.clone()
1238    }
1239}
1240
1241#[cfg(test)]
1242mod test {
1243    use std::sync::Arc;
1244
1245    use arrow::array::{
1246        BooleanArray, TimestampMicrosecondArray, TimestampMillisecondArray,
1247        TimestampNanosecondArray, TimestampSecondArray,
1248    };
1249    use arrow::json::ArrayWriter;
1250    use arrow_schema::{DataType, Field, Schema, SortOptions, TimeUnit};
1251    use common_time::Timestamp;
1252    use datafusion_physical_expr::expressions::Column;
1253    use futures::StreamExt;
1254    use store_api::region_engine::PartitionRange;
1255
1256    use super::*;
1257    use crate::test_util::{MockInputExec, new_ts_array};
1258
1259    #[ignore = "hard to gen expected data correctly here, TODO(discord9): fix it later"]
1260    #[tokio::test]
1261    async fn fuzzy_test() {
1262        let test_cnt = 100;
1263        // bound for total count of PartitionRange
1264        let part_cnt_bound = 100;
1265        // bound for timestamp range size and offset for each PartitionRange
1266        let range_size_bound = 100;
1267        let range_offset_bound = 100;
1268        // bound for batch count and size within each PartitionRange
1269        let batch_cnt_bound = 20;
1270        let batch_size_bound = 100;
1271
1272        let mut rng = fastrand::Rng::new();
1273        rng.seed(1337);
1274
1275        let mut test_cases = Vec::new();
1276
1277        for case_id in 0..test_cnt {
1278            let mut bound_val: Option<i64> = None;
1279            let descending = rng.bool();
1280            let nulls_first = rng.bool();
1281            let opt = SortOptions {
1282                descending,
1283                nulls_first,
1284            };
1285            let limit = if rng.bool() {
1286                Some(rng.usize(1..batch_cnt_bound * batch_size_bound))
1287            } else {
1288                None
1289            };
1290            let unit = match rng.u8(0..3) {
1291                0 => TimeUnit::Second,
1292                1 => TimeUnit::Millisecond,
1293                2 => TimeUnit::Microsecond,
1294                _ => TimeUnit::Nanosecond,
1295            };
1296
1297            let schema = Schema::new(vec![Field::new(
1298                "ts",
1299                DataType::Timestamp(unit, None),
1300                false,
1301            )]);
1302            let schema = Arc::new(schema);
1303
1304            let mut input_ranged_data = vec![];
1305            let mut output_ranges = vec![];
1306            let mut output_data = vec![];
1307            // generate each input `PartitionRange`
1308            for part_id in 0..rng.usize(0..part_cnt_bound) {
1309                // generate each `PartitionRange`'s timestamp range
1310                let (start, end) = if descending {
1311                    // Use 1..=range_offset_bound to ensure strictly decreasing end values
1312                    let end = bound_val
1313                        .map(
1314                            |i| i
1315                            .checked_sub(rng.i64(1..=range_offset_bound))
1316                            .expect("Bad luck, fuzzy test generate data that will overflow, change seed and try again")
1317                        )
1318                        .unwrap_or_else(|| rng.i64(-100000000..100000000));
1319                    bound_val = Some(end);
1320                    let start = end - rng.i64(1..range_size_bound);
1321                    let start = Timestamp::new(start, unit.into());
1322                    let end = Timestamp::new(end, unit.into());
1323                    (start, end)
1324                } else {
1325                    // Use 1..=range_offset_bound to ensure strictly increasing start values
1326                    let start = bound_val
1327                        .map(|i| i + rng.i64(1..=range_offset_bound))
1328                        .unwrap_or_else(|| rng.i64(..));
1329                    bound_val = Some(start);
1330                    let end = start + rng.i64(1..range_size_bound);
1331                    let start = Timestamp::new(start, unit.into());
1332                    let end = Timestamp::new(end, unit.into());
1333                    (start, end)
1334                };
1335                assert!(start < end);
1336
1337                let mut per_part_sort_data = vec![];
1338                let mut batches = vec![];
1339                for _batch_idx in 0..rng.usize(1..batch_cnt_bound) {
1340                    let cnt = rng.usize(0..batch_size_bound) + 1;
1341                    let iter = 0..rng.usize(0..cnt);
1342                    let mut data_gen = iter
1343                        .map(|_| rng.i64(start.value()..end.value()))
1344                        .collect_vec();
1345                    if data_gen.is_empty() {
1346                        // current batch is empty, skip
1347                        continue;
1348                    }
1349                    // mito always sort on ASC order
1350                    data_gen.sort();
1351                    per_part_sort_data.extend(data_gen.clone());
1352                    let arr = new_ts_array(unit, data_gen.clone());
1353                    let batch = DfRecordBatch::try_new(schema.clone(), vec![arr]).unwrap();
1354                    batches.push(batch);
1355                }
1356
1357                let range = PartitionRange {
1358                    start,
1359                    end,
1360                    num_rows: batches.iter().map(|b| b.num_rows()).sum(),
1361                    identifier: part_id,
1362                };
1363                input_ranged_data.push((range, batches));
1364
1365                output_ranges.push(range);
1366                if per_part_sort_data.is_empty() {
1367                    continue;
1368                }
1369                output_data.extend_from_slice(&per_part_sort_data);
1370            }
1371
1372            // adjust output data with adjacent PartitionRanges
1373            let mut output_data_iter = output_data.iter().peekable();
1374            let mut output_data = vec![];
1375            for range in output_ranges.clone() {
1376                let mut cur_data = vec![];
1377                while let Some(val) = output_data_iter.peek() {
1378                    if **val < range.start.value() || **val >= range.end.value() {
1379                        break;
1380                    }
1381                    cur_data.push(*output_data_iter.next().unwrap());
1382                }
1383
1384                if cur_data.is_empty() {
1385                    continue;
1386                }
1387
1388                if descending {
1389                    cur_data.sort_by(|a, b| b.cmp(a));
1390                } else {
1391                    cur_data.sort();
1392                }
1393                output_data.push(cur_data);
1394            }
1395
1396            let expected_output = if let Some(limit) = limit {
1397                let mut accumulated = Vec::new();
1398                let mut seen = 0usize;
1399                for mut range_values in output_data {
1400                    seen += range_values.len();
1401                    accumulated.append(&mut range_values);
1402                    if seen >= limit {
1403                        break;
1404                    }
1405                }
1406
1407                if accumulated.is_empty() {
1408                    None
1409                } else {
1410                    if descending {
1411                        accumulated.sort_by(|a, b| b.cmp(a));
1412                    } else {
1413                        accumulated.sort();
1414                    }
1415                    accumulated.truncate(limit.min(accumulated.len()));
1416
1417                    Some(
1418                        DfRecordBatch::try_new(
1419                            schema.clone(),
1420                            vec![new_ts_array(unit, accumulated)],
1421                        )
1422                        .unwrap(),
1423                    )
1424                }
1425            } else {
1426                let batches = output_data
1427                    .into_iter()
1428                    .map(|a| {
1429                        DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, a)]).unwrap()
1430                    })
1431                    .collect_vec();
1432                if batches.is_empty() {
1433                    None
1434                } else {
1435                    Some(concat_batches(&schema, &batches).unwrap())
1436                }
1437            };
1438
1439            test_cases.push((
1440                case_id,
1441                unit,
1442                input_ranged_data,
1443                schema,
1444                opt,
1445                limit,
1446                expected_output,
1447            ));
1448        }
1449
1450        for (case_id, _unit, input_ranged_data, schema, opt, limit, expected_output) in test_cases {
1451            run_test(
1452                case_id,
1453                input_ranged_data,
1454                schema,
1455                opt,
1456                limit,
1457                expected_output,
1458                None,
1459            )
1460            .await;
1461        }
1462    }
1463
1464    #[tokio::test]
1465    async fn simple_cases() {
1466        let testcases = vec![
1467            (
1468                TimeUnit::Millisecond,
1469                vec![
1470                    ((0, 10), vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]]),
1471                    ((5, 10), vec![vec![5, 6], vec![7, 8]]),
1472                ],
1473                false,
1474                None,
1475                vec![vec![1, 2, 3, 4, 5, 5, 6, 6, 7, 7, 8, 8, 9]],
1476            ),
1477            // Case 1: Descending sort with overlapping ranges that have the same primary end (end=10).
1478            // Ranges [5,10) and [0,10) are grouped together, so their data is merged before sorting.
1479            (
1480                TimeUnit::Millisecond,
1481                vec![
1482                    ((5, 10), vec![vec![5, 6], vec![7, 8, 9]]),
1483                    ((0, 10), vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8]]),
1484                ],
1485                true,
1486                None,
1487                vec![vec![9, 8, 8, 7, 7, 6, 6, 5, 5, 4, 3, 2, 1]],
1488            ),
1489            (
1490                TimeUnit::Millisecond,
1491                vec![
1492                    ((5, 10), vec![]),
1493                    ((0, 10), vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8]]),
1494                ],
1495                true,
1496                None,
1497                vec![vec![8, 7, 6, 5, 4, 3, 2, 1]],
1498            ),
1499            (
1500                TimeUnit::Millisecond,
1501                vec![
1502                    ((15, 20), vec![vec![17, 18, 19]]),
1503                    ((10, 15), vec![]),
1504                    ((5, 10), vec![]),
1505                    ((0, 10), vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8]]),
1506                ],
1507                true,
1508                None,
1509                vec![vec![19, 18, 17], vec![8, 7, 6, 5, 4, 3, 2, 1]],
1510            ),
1511            (
1512                TimeUnit::Millisecond,
1513                vec![
1514                    ((15, 20), vec![]),
1515                    ((10, 15), vec![]),
1516                    ((5, 10), vec![]),
1517                    ((0, 10), vec![]),
1518                ],
1519                true,
1520                None,
1521                vec![],
1522            ),
1523            // Case 5: Data from one batch spans multiple ranges. Ranges with same end are grouped.
1524            // Ranges: [15,20) end=20, [10,15) end=15, [5,10) end=10, [0,10) end=10
1525            // Groups: {[15,20)}, {[10,15)}, {[5,10), [0,10)}
1526            // The last two ranges are merged because they share end=10.
1527            (
1528                TimeUnit::Millisecond,
1529                vec![
1530                    (
1531                        (15, 20),
1532                        vec![vec![15, 17, 19, 10, 11, 12, 5, 6, 7, 8, 9, 1, 2, 3, 4]],
1533                    ),
1534                    ((10, 15), vec![]),
1535                    ((5, 10), vec![]),
1536                    ((0, 10), vec![]),
1537                ],
1538                true,
1539                None,
1540                vec![
1541                    vec![19, 17, 15],
1542                    vec![12, 11, 10],
1543                    vec![9, 8, 7, 6, 5, 4, 3, 2, 1],
1544                ],
1545            ),
1546            (
1547                TimeUnit::Millisecond,
1548                vec![
1549                    (
1550                        (15, 20),
1551                        vec![vec![15, 17, 19, 10, 11, 12, 5, 6, 7, 8, 9, 1, 2, 3, 4]],
1552                    ),
1553                    ((10, 15), vec![]),
1554                    ((5, 10), vec![]),
1555                    ((0, 10), vec![]),
1556                ],
1557                true,
1558                Some(2),
1559                vec![vec![19, 17]],
1560            ),
1561        ];
1562
1563        for (identifier, (unit, input_ranged_data, descending, limit, expected_output)) in
1564            testcases.into_iter().enumerate()
1565        {
1566            let schema = Schema::new(vec![Field::new(
1567                "ts",
1568                DataType::Timestamp(unit, None),
1569                false,
1570            )]);
1571            let schema = Arc::new(schema);
1572            let opt = SortOptions {
1573                descending,
1574                ..Default::default()
1575            };
1576
1577            let input_ranged_data = input_ranged_data
1578                .into_iter()
1579                .map(|(range, data)| {
1580                    let part = PartitionRange {
1581                        start: Timestamp::new(range.0, unit.into()),
1582                        end: Timestamp::new(range.1, unit.into()),
1583                        num_rows: data.iter().map(|b| b.len()).sum(),
1584                        identifier,
1585                    };
1586
1587                    let batches = data
1588                        .into_iter()
1589                        .map(|b| {
1590                            let arr = new_ts_array(unit, b);
1591                            DfRecordBatch::try_new(schema.clone(), vec![arr]).unwrap()
1592                        })
1593                        .collect_vec();
1594                    (part, batches)
1595                })
1596                .collect_vec();
1597
1598            let expected_output = expected_output
1599                .into_iter()
1600                .map(|a| {
1601                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, a)]).unwrap()
1602                })
1603                .collect_vec();
1604            let expected_output = if expected_output.is_empty() {
1605                None
1606            } else {
1607                Some(concat_batches(&schema, &expected_output).unwrap())
1608            };
1609
1610            run_test(
1611                identifier,
1612                input_ranged_data,
1613                schema.clone(),
1614                opt,
1615                limit,
1616                expected_output,
1617                None,
1618            )
1619            .await;
1620        }
1621    }
1622
1623    #[allow(clippy::print_stdout)]
1624    async fn run_test(
1625        case_id: usize,
1626        input_ranged_data: Vec<(PartitionRange, Vec<DfRecordBatch>)>,
1627        schema: SchemaRef,
1628        opt: SortOptions,
1629        limit: Option<usize>,
1630        expected_output: Option<DfRecordBatch>,
1631        expected_polled_rows: Option<usize>,
1632    ) {
1633        if let (Some(limit), Some(rb)) = (limit, &expected_output) {
1634            assert!(
1635                rb.num_rows() <= limit,
1636                "Expect row count in expected output({}) <= limit({})",
1637                rb.num_rows(),
1638                limit
1639            );
1640        }
1641
1642        let mut data_partition = Vec::with_capacity(input_ranged_data.len());
1643        let mut ranges = Vec::with_capacity(input_ranged_data.len());
1644        for (part_range, batches) in input_ranged_data {
1645            data_partition.push(batches);
1646            ranges.push(part_range);
1647        }
1648
1649        let mock_input = Arc::new(MockInputExec::new(data_partition, schema.clone()));
1650
1651        let exec = PartSortExec::try_new(
1652            PhysicalSortExpr {
1653                expr: Arc::new(Column::new("ts", 0)),
1654                options: opt,
1655            },
1656            limit,
1657            vec![ranges.clone()],
1658            mock_input.clone(),
1659        )
1660        .unwrap();
1661
1662        let exec_stream = exec.execute(0, Arc::new(TaskContext::default())).unwrap();
1663
1664        let real_output = exec_stream.map(|r| r.unwrap()).collect::<Vec<_>>().await;
1665        if limit.is_some() {
1666            assert!(
1667                real_output.len() <= 1,
1668                "case_{case_id} expects a single output batch when limit is set, got {}",
1669                real_output.len()
1670            );
1671        }
1672
1673        let actual_output = if real_output.is_empty() {
1674            None
1675        } else {
1676            Some(concat_batches(&schema, &real_output).unwrap())
1677        };
1678
1679        if let Some(expected_polled_rows) = expected_polled_rows {
1680            let input_pulled_rows = mock_input.metrics().unwrap().output_rows().unwrap();
1681            assert_eq!(input_pulled_rows, expected_polled_rows);
1682        }
1683
1684        match (actual_output, expected_output) {
1685            (None, None) => {}
1686            (Some(actual), Some(expected)) => {
1687                if actual != expected {
1688                    let mut actual_json: Vec<u8> = Vec::new();
1689                    let mut writer = ArrayWriter::new(&mut actual_json);
1690                    writer.write(&actual).unwrap();
1691                    writer.finish().unwrap();
1692
1693                    let mut expected_json: Vec<u8> = Vec::new();
1694                    let mut writer = ArrayWriter::new(&mut expected_json);
1695                    writer.write(&expected).unwrap();
1696                    writer.finish().unwrap();
1697
1698                    panic!(
1699                        "case_{} failed (limit {limit:?}), opt: {:?},\nreal_output: {}\nexpected: {}",
1700                        case_id,
1701                        opt,
1702                        String::from_utf8_lossy(&actual_json),
1703                        String::from_utf8_lossy(&expected_json),
1704                    );
1705                }
1706            }
1707            (None, Some(expected)) => panic!(
1708                "case_{} failed (limit {limit:?}), opt: {:?},\nreal output is empty, expected {} rows",
1709                case_id,
1710                opt,
1711                expected.num_rows()
1712            ),
1713            (Some(actual), None) => panic!(
1714                "case_{} failed (limit {limit:?}), opt: {:?},\nreal output has {} rows, expected empty",
1715                case_id,
1716                opt,
1717                actual.num_rows()
1718            ),
1719        }
1720    }
1721
1722    /// Test that verifies the limit is correctly applied per partition when
1723    /// multiple batches are received for the same partition.
1724    #[tokio::test]
1725    async fn test_limit_with_multiple_batches_per_partition() {
1726        let unit = TimeUnit::Millisecond;
1727        let schema = Arc::new(Schema::new(vec![Field::new(
1728            "ts",
1729            DataType::Timestamp(unit, None),
1730            false,
1731        )]));
1732
1733        // Test case: Multiple batches in a single partition with limit=3
1734        // Input: 3 batches with [1,2,3], [4,5,6], [7,8,9] all in partition (0,10)
1735        // Expected: Only top 3 values [9,8,7] for descending sort
1736        let input_ranged_data = vec![(
1737            PartitionRange {
1738                start: Timestamp::new(0, unit.into()),
1739                end: Timestamp::new(10, unit.into()),
1740                num_rows: 9,
1741                identifier: 0,
1742            },
1743            vec![
1744                DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![1, 2, 3])])
1745                    .unwrap(),
1746                DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![4, 5, 6])])
1747                    .unwrap(),
1748                DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![7, 8, 9])])
1749                    .unwrap(),
1750            ],
1751        )];
1752
1753        let expected_output = Some(
1754            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![9, 8, 7])])
1755                .unwrap(),
1756        );
1757
1758        run_test(
1759            1000,
1760            input_ranged_data,
1761            schema.clone(),
1762            SortOptions {
1763                descending: true,
1764                ..Default::default()
1765            },
1766            Some(3),
1767            expected_output,
1768            None,
1769        )
1770        .await;
1771
1772        // Test case: Multiple batches across multiple partitions with limit=2
1773        // Partition 0: batches [10,11,12], [13,14,15] -> top 2 descending = [15,14]
1774        // Partition 1: batches [1,2,3], [4,5] -> top 2 descending = [5,4]
1775        let input_ranged_data = vec![
1776            (
1777                PartitionRange {
1778                    start: Timestamp::new(10, unit.into()),
1779                    end: Timestamp::new(20, unit.into()),
1780                    num_rows: 6,
1781                    identifier: 0,
1782                },
1783                vec![
1784                    DfRecordBatch::try_new(
1785                        schema.clone(),
1786                        vec![new_ts_array(unit, vec![10, 11, 12])],
1787                    )
1788                    .unwrap(),
1789                    DfRecordBatch::try_new(
1790                        schema.clone(),
1791                        vec![new_ts_array(unit, vec![13, 14, 15])],
1792                    )
1793                    .unwrap(),
1794                ],
1795            ),
1796            (
1797                PartitionRange {
1798                    start: Timestamp::new(0, unit.into()),
1799                    end: Timestamp::new(10, unit.into()),
1800                    num_rows: 5,
1801                    identifier: 1,
1802                },
1803                vec![
1804                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![1, 2, 3])])
1805                        .unwrap(),
1806                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![4, 5])])
1807                        .unwrap(),
1808                ],
1809            ),
1810        ];
1811
1812        let expected_output = Some(
1813            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![15, 14])]).unwrap(),
1814        );
1815
1816        run_test(
1817            1001,
1818            input_ranged_data,
1819            schema.clone(),
1820            SortOptions {
1821                descending: true,
1822                ..Default::default()
1823            },
1824            Some(2),
1825            expected_output,
1826            None,
1827        )
1828        .await;
1829
1830        // Test case: Ascending sort with limit
1831        // Partition: batches [7,8,9], [4,5,6], [1,2,3] -> top 2 ascending = [1,2]
1832        let input_ranged_data = vec![(
1833            PartitionRange {
1834                start: Timestamp::new(0, unit.into()),
1835                end: Timestamp::new(10, unit.into()),
1836                num_rows: 9,
1837                identifier: 0,
1838            },
1839            vec![
1840                DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![7, 8, 9])])
1841                    .unwrap(),
1842                DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![4, 5, 6])])
1843                    .unwrap(),
1844                DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![1, 2, 3])])
1845                    .unwrap(),
1846            ],
1847        )];
1848
1849        let expected_output = Some(
1850            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![1, 2])]).unwrap(),
1851        );
1852
1853        run_test(
1854            1002,
1855            input_ranged_data,
1856            schema.clone(),
1857            SortOptions {
1858                descending: false,
1859                ..Default::default()
1860            },
1861            Some(2),
1862            expected_output,
1863            None,
1864        )
1865        .await;
1866    }
1867
1868    #[test]
1869    fn dynamic_expressions_produced_returns_topk_filter_arc() {
1870        let unit = TimeUnit::Millisecond;
1871        let schema = Arc::new(Schema::new(vec![Field::new(
1872            "ts",
1873            DataType::Timestamp(unit, None),
1874            false,
1875        )]));
1876        let partition_range = PartitionRange {
1877            start: Timestamp::new(0, unit.into()),
1878            end: Timestamp::new(10, unit.into()),
1879            num_rows: 0,
1880            identifier: 0,
1881        };
1882        let sort_expr = PhysicalSortExpr {
1883            expr: Arc::new(Column::new("ts", 0)),
1884            options: SortOptions::default(),
1885        };
1886
1887        let limited = PartSortExec::try_new(
1888            sort_expr.clone(),
1889            Some(1),
1890            vec![vec![partition_range]],
1891            Arc::new(MockInputExec::new(vec![vec![]], schema.clone())),
1892        )
1893        .unwrap();
1894        let expected = limited.dynamic_filter.as_ref().unwrap().clone() as Arc<dyn PhysicalExpr>;
1895        let produced = limited.dynamic_expressions_produced();
1896        assert_eq!(produced.len(), 1);
1897        assert!(Arc::ptr_eq(&produced[0], &expected));
1898
1899        let mut applied_dynamic_filter = None;
1900        limited
1901            .apply_expressions(&mut |expr| {
1902                if expr.expression_id().is_some() {
1903                    applied_dynamic_filter = Some(expr.clone());
1904                }
1905                Ok(TreeNodeRecursion::Continue)
1906            })
1907            .unwrap();
1908        let applied_dynamic_filter = applied_dynamic_filter.unwrap();
1909        assert!(Arc::ptr_eq(&produced[0], &applied_dynamic_filter));
1910        assert_eq!(
1911            produced[0].expression_id(),
1912            applied_dynamic_filter.expression_id()
1913        );
1914
1915        let unlimited = PartSortExec::try_new(
1916            sort_expr,
1917            None,
1918            vec![vec![partition_range]],
1919            Arc::new(MockInputExec::new(vec![vec![]], schema)),
1920        )
1921        .unwrap();
1922        assert!(unlimited.dynamic_expressions_produced().is_empty());
1923    }
1924
1925    #[test]
1926    fn test_topk_buffer_is_bounded_and_updates_dynamic_filter() {
1927        let unit = TimeUnit::Millisecond;
1928        let schema = Arc::new(Schema::new(vec![Field::new(
1929            "ts",
1930            DataType::Timestamp(unit, None),
1931            false,
1932        )]));
1933        let sort_data_type = DataType::Timestamp(unit, None);
1934        let partition_range = PartitionRange {
1935            start: Timestamp::new(0, unit.into()),
1936            end: Timestamp::new(10, unit.into()),
1937            num_rows: 9,
1938            identifier: 0,
1939        };
1940        let mock_input = Arc::new(MockInputExec::new(vec![vec![]], schema.clone()));
1941        let exec = PartSortExec::try_new(
1942            PhysicalSortExpr {
1943                expr: Arc::new(Column::new("ts", 0)),
1944                options: SortOptions {
1945                    descending: true,
1946                    ..Default::default()
1947                },
1948            },
1949            Some(3),
1950            vec![vec![partition_range]],
1951            mock_input.clone(),
1952        )
1953        .unwrap();
1954        let input_stream = mock_input
1955            .execute(0, Arc::new(TaskContext::default()))
1956            .unwrap();
1957        let mut stream = PartSortStream::new(
1958            Arc::new(TaskContext::default()),
1959            &exec,
1960            Some(3),
1961            input_stream,
1962            vec![partition_range],
1963            0,
1964        )
1965        .unwrap();
1966
1967        for batch in [vec![1, 2, 3], vec![4, 5, 6], vec![0, 7, 8]] {
1968            stream
1969                .push_buffer(
1970                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, batch)])
1971                        .unwrap(),
1972                    &sort_data_type,
1973                )
1974                .unwrap();
1975            assert_eq!(stream.buffer.num_rows(), 3);
1976        }
1977
1978        let dynamic_filter = stream.dynamic_filter.as_ref().unwrap().clone();
1979        let probe = DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![5, 6, 7])])
1980            .unwrap();
1981        let predicate = dynamic_filter.current().unwrap();
1982        let result = predicate
1983            .evaluate(&probe)
1984            .unwrap()
1985            .into_array(probe.num_rows())
1986            .unwrap();
1987        let result = result.as_any().downcast_ref::<BooleanArray>().unwrap();
1988        assert_eq!(result, &BooleanArray::from(vec![false, false, true]));
1989
1990        let expected =
1991            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![8, 7, 6])])
1992                .unwrap();
1993        assert_eq!(stream.sort_buffer().unwrap(), expected);
1994    }
1995
1996    #[test]
1997    fn test_topk_limit_zero_clears_buffer_without_threshold() {
1998        let unit = TimeUnit::Millisecond;
1999        let schema = Arc::new(Schema::new(vec![Field::new(
2000            "ts",
2001            DataType::Timestamp(unit, None),
2002            false,
2003        )]));
2004        let sort_data_type = DataType::Timestamp(unit, None);
2005        let partition_range = PartitionRange {
2006            start: Timestamp::new(0, unit.into()),
2007            end: Timestamp::new(10, unit.into()),
2008            num_rows: 3,
2009            identifier: 0,
2010        };
2011        let mock_input = Arc::new(MockInputExec::new(vec![vec![]], schema.clone()));
2012        let exec = PartSortExec::try_new(
2013            PhysicalSortExpr {
2014                expr: Arc::new(Column::new("ts", 0)),
2015                options: SortOptions {
2016                    descending: true,
2017                    ..Default::default()
2018                },
2019            },
2020            Some(0),
2021            vec![vec![partition_range]],
2022            mock_input.clone(),
2023        )
2024        .unwrap();
2025        let input_stream = mock_input
2026            .execute(0, Arc::new(TaskContext::default()))
2027            .unwrap();
2028        let mut stream = PartSortStream::new(
2029            Arc::new(TaskContext::default()),
2030            &exec,
2031            Some(0),
2032            input_stream,
2033            vec![partition_range],
2034            0,
2035        )
2036        .unwrap();
2037
2038        stream
2039            .push_buffer(
2040                DfRecordBatch::try_new(schema, vec![new_ts_array(unit, vec![1, 2, 3])]).unwrap(),
2041                &sort_data_type,
2042            )
2043            .unwrap();
2044
2045        assert_eq!(stream.buffer.num_rows(), 0);
2046        assert_eq!(stream.dynamic_filter_threshold, None);
2047    }
2048
2049    /// Test that verifies early termination behavior.
2050    /// Once we've produced limit * num_partitions rows, we should stop
2051    /// pulling from input stream.
2052    #[tokio::test]
2053    async fn test_early_termination() {
2054        let unit = TimeUnit::Millisecond;
2055        let schema = Arc::new(Schema::new(vec![Field::new(
2056            "ts",
2057            DataType::Timestamp(unit, None),
2058            false,
2059        )]));
2060
2061        // Create 3 partitions, each with more data than the limit
2062        // limit=2 per partition, so total expected output = 6 rows
2063        // After producing 6 rows, early termination should kick in
2064        // For descending sort, ranges must be ordered by (end DESC, start DESC)
2065        let input_ranged_data = vec![
2066            (
2067                PartitionRange {
2068                    start: Timestamp::new(20, unit.into()),
2069                    end: Timestamp::new(30, unit.into()),
2070                    num_rows: 10,
2071                    identifier: 2,
2072                },
2073                vec![
2074                    DfRecordBatch::try_new(
2075                        schema.clone(),
2076                        vec![new_ts_array(unit, vec![21, 22, 23, 24, 25])],
2077                    )
2078                    .unwrap(),
2079                    DfRecordBatch::try_new(
2080                        schema.clone(),
2081                        vec![new_ts_array(unit, vec![26, 27, 28, 29, 30])],
2082                    )
2083                    .unwrap(),
2084                ],
2085            ),
2086            (
2087                PartitionRange {
2088                    start: Timestamp::new(10, unit.into()),
2089                    end: Timestamp::new(20, unit.into()),
2090                    num_rows: 10,
2091                    identifier: 1,
2092                },
2093                vec![
2094                    DfRecordBatch::try_new(
2095                        schema.clone(),
2096                        vec![new_ts_array(unit, vec![11, 12, 13, 14, 15])],
2097                    )
2098                    .unwrap(),
2099                    DfRecordBatch::try_new(
2100                        schema.clone(),
2101                        vec![new_ts_array(unit, vec![16, 17, 18, 19, 20])],
2102                    )
2103                    .unwrap(),
2104                ],
2105            ),
2106            (
2107                PartitionRange {
2108                    start: Timestamp::new(0, unit.into()),
2109                    end: Timestamp::new(10, unit.into()),
2110                    num_rows: 10,
2111                    identifier: 0,
2112                },
2113                vec![
2114                    DfRecordBatch::try_new(
2115                        schema.clone(),
2116                        vec![new_ts_array(unit, vec![1, 2, 3, 4, 5])],
2117                    )
2118                    .unwrap(),
2119                    DfRecordBatch::try_new(
2120                        schema.clone(),
2121                        vec![new_ts_array(unit, vec![6, 7, 8, 9, 10])],
2122                    )
2123                    .unwrap(),
2124                ],
2125            ),
2126        ];
2127
2128        // PartSort won't reorder `PartitionRange` (it assumes it's already ordered), so it will not read other partitions.
2129        // This case is just to verify that early termination works as expected.
2130        // First partition [20, 30) produces top 2 values: 29, 28
2131        let expected_output = Some(
2132            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![29, 28])]).unwrap(),
2133        );
2134
2135        run_test(
2136            1003,
2137            input_ranged_data,
2138            schema.clone(),
2139            SortOptions {
2140                descending: true,
2141                ..Default::default()
2142            },
2143            Some(2),
2144            expected_output,
2145            Some(10),
2146        )
2147        .await;
2148    }
2149
2150    /// Example:
2151    /// - Range [70, 100) has data [80, 90, 95]
2152    /// - Range [50, 100) has data [55, 65, 75, 85, 95]
2153    #[tokio::test]
2154    async fn test_primary_end_grouping_with_limit() {
2155        let unit = TimeUnit::Millisecond;
2156        let schema = Arc::new(Schema::new(vec![Field::new(
2157            "ts",
2158            DataType::Timestamp(unit, None),
2159            false,
2160        )]));
2161
2162        // Two ranges with the same end (100) - they should be grouped together
2163        // For descending, ranges are ordered by (end DESC, start DESC)
2164        // So [70, 100) comes before [50, 100) (70 > 50)
2165        let input_ranged_data = vec![
2166            (
2167                PartitionRange {
2168                    start: Timestamp::new(70, unit.into()),
2169                    end: Timestamp::new(100, unit.into()),
2170                    num_rows: 3,
2171                    identifier: 0,
2172                },
2173                vec![
2174                    DfRecordBatch::try_new(
2175                        schema.clone(),
2176                        vec![new_ts_array(unit, vec![80, 90, 95])],
2177                    )
2178                    .unwrap(),
2179                ],
2180            ),
2181            (
2182                PartitionRange {
2183                    start: Timestamp::new(50, unit.into()),
2184                    end: Timestamp::new(100, unit.into()),
2185                    num_rows: 5,
2186                    identifier: 1,
2187                },
2188                vec![
2189                    DfRecordBatch::try_new(
2190                        schema.clone(),
2191                        vec![new_ts_array(unit, vec![55, 65, 75, 85, 95])],
2192                    )
2193                    .unwrap(),
2194                ],
2195            ),
2196        ];
2197
2198        // With limit=4, descending: top 4 values from combined data
2199        // Combined: [80, 90, 95, 55, 65, 75, 85, 95] -> sorted desc: [95, 95, 90, 85, 80, 75, 65, 55]
2200        // Top 4: [95, 95, 90, 85]
2201        let expected_output = Some(
2202            DfRecordBatch::try_new(
2203                schema.clone(),
2204                vec![new_ts_array(unit, vec![95, 95, 90, 85])],
2205            )
2206            .unwrap(),
2207        );
2208
2209        run_test(
2210            2000,
2211            input_ranged_data,
2212            schema.clone(),
2213            SortOptions {
2214                descending: true,
2215                ..Default::default()
2216            },
2217            Some(4),
2218            expected_output,
2219            None,
2220        )
2221        .await;
2222    }
2223
2224    /// Test case with three ranges demonstrating the "keep pulling" behavior.
2225    /// After processing ranges with end=100, the smallest value in top-k might still
2226    /// be reachable by the next group.
2227    ///
2228    /// Ranges: [70, 100), [50, 100), [40, 95)
2229    /// With descending sort and limit=4:
2230    /// - Group 1 (end=100): [70, 100) and [50, 100) merged
2231    /// - Group 2 (end=95): [40, 95)
2232    /// After group 1, smallest in top-4 is 85. Range [40, 95) could have values >= 85,
2233    /// so we continue to group 2.
2234    #[tokio::test]
2235    async fn test_three_ranges_keep_pulling() {
2236        let unit = TimeUnit::Millisecond;
2237        let schema = Arc::new(Schema::new(vec![Field::new(
2238            "ts",
2239            DataType::Timestamp(unit, None),
2240            false,
2241        )]));
2242
2243        // Three ranges, two with same end (100), one with different end (95)
2244        let input_ranged_data = vec![
2245            (
2246                PartitionRange {
2247                    start: Timestamp::new(70, unit.into()),
2248                    end: Timestamp::new(100, unit.into()),
2249                    num_rows: 3,
2250                    identifier: 0,
2251                },
2252                vec![
2253                    DfRecordBatch::try_new(
2254                        schema.clone(),
2255                        vec![new_ts_array(unit, vec![80, 90, 95])],
2256                    )
2257                    .unwrap(),
2258                ],
2259            ),
2260            (
2261                PartitionRange {
2262                    start: Timestamp::new(50, unit.into()),
2263                    end: Timestamp::new(100, unit.into()),
2264                    num_rows: 3,
2265                    identifier: 1,
2266                },
2267                vec![
2268                    DfRecordBatch::try_new(
2269                        schema.clone(),
2270                        vec![new_ts_array(unit, vec![55, 75, 85])],
2271                    )
2272                    .unwrap(),
2273                ],
2274            ),
2275            (
2276                PartitionRange {
2277                    start: Timestamp::new(40, unit.into()),
2278                    end: Timestamp::new(95, unit.into()),
2279                    num_rows: 3,
2280                    identifier: 2,
2281                },
2282                vec![
2283                    DfRecordBatch::try_new(
2284                        schema.clone(),
2285                        vec![new_ts_array(unit, vec![45, 65, 94])],
2286                    )
2287                    .unwrap(),
2288                ],
2289            ),
2290        ];
2291
2292        // All data: [80, 90, 95, 55, 75, 85, 45, 65, 94]
2293        // Sorted descending: [95, 94, 90, 85, 80, 75, 65, 55, 45]
2294        // With limit=4: should be top 4 largest values across all ranges: [95, 94, 90, 85]
2295        let expected_output = Some(
2296            DfRecordBatch::try_new(
2297                schema.clone(),
2298                vec![new_ts_array(unit, vec![95, 94, 90, 85])],
2299            )
2300            .unwrap(),
2301        );
2302
2303        run_test(
2304            2001,
2305            input_ranged_data,
2306            schema.clone(),
2307            SortOptions {
2308                descending: true,
2309                ..Default::default()
2310            },
2311            Some(4),
2312            expected_output,
2313            None,
2314        )
2315        .await;
2316    }
2317
2318    /// Test early termination based on threshold comparison with next group.
2319    /// When the threshold (smallest value for descending) is >= next group's primary end,
2320    /// we can stop early because the next group cannot have better values.
2321    #[tokio::test]
2322    async fn test_threshold_based_early_termination() {
2323        let unit = TimeUnit::Millisecond;
2324        let schema = Arc::new(Schema::new(vec![Field::new(
2325            "ts",
2326            DataType::Timestamp(unit, None),
2327            false,
2328        )]));
2329
2330        // Group 1 (end=100) has 6 rows, TopK will keep top 4
2331        // Group 2 (end=90) has 3 rows - should NOT be processed because
2332        // threshold (96) >= next_primary_end (90)
2333        let input_ranged_data = vec![
2334            (
2335                PartitionRange {
2336                    start: Timestamp::new(70, unit.into()),
2337                    end: Timestamp::new(100, unit.into()),
2338                    num_rows: 6,
2339                    identifier: 0,
2340                },
2341                vec![
2342                    DfRecordBatch::try_new(
2343                        schema.clone(),
2344                        vec![new_ts_array(unit, vec![94, 95, 96, 97, 98, 99])],
2345                    )
2346                    .unwrap(),
2347                ],
2348            ),
2349            (
2350                PartitionRange {
2351                    start: Timestamp::new(50, unit.into()),
2352                    end: Timestamp::new(90, unit.into()),
2353                    num_rows: 3,
2354                    identifier: 1,
2355                },
2356                vec![
2357                    DfRecordBatch::try_new(
2358                        schema.clone(),
2359                        vec![new_ts_array(unit, vec![85, 86, 87])],
2360                    )
2361                    .unwrap(),
2362                ],
2363            ),
2364        ];
2365
2366        // With limit=4, descending: top 4 from group 1 are [99, 98, 97, 96]
2367        // Threshold is 96, next group's primary_end is 90
2368        // Since 96 >= 90, we stop after group 1
2369        let expected_output = Some(
2370            DfRecordBatch::try_new(
2371                schema.clone(),
2372                vec![new_ts_array(unit, vec![99, 98, 97, 96])],
2373            )
2374            .unwrap(),
2375        );
2376
2377        run_test(
2378            2002,
2379            input_ranged_data,
2380            schema.clone(),
2381            SortOptions {
2382                descending: true,
2383                ..Default::default()
2384            },
2385            Some(4),
2386            expected_output,
2387            Some(9), // Pull both batches since all rows fall within the first range
2388        )
2389        .await;
2390    }
2391
2392    /// Test that we continue to next group when threshold is within next group's range.
2393    /// Even after fulfilling limit, if threshold < next_primary_end (descending),
2394    /// we would need to continue... but limit exhaustion stops us first.
2395    #[tokio::test]
2396    async fn test_continue_when_threshold_in_next_group_range() {
2397        let unit = TimeUnit::Millisecond;
2398        let schema = Arc::new(Schema::new(vec![Field::new(
2399            "ts",
2400            DataType::Timestamp(unit, None),
2401            false,
2402        )]));
2403
2404        // Group 1 (end=100) has 6 rows, TopK will keep top 4
2405        // Group 2 (end=98) has 3 rows - threshold (96) < 98, so next group
2406        // could theoretically have better values. Continue reading.
2407        let input_ranged_data = vec![
2408            (
2409                PartitionRange {
2410                    start: Timestamp::new(90, unit.into()),
2411                    end: Timestamp::new(100, unit.into()),
2412                    num_rows: 6,
2413                    identifier: 0,
2414                },
2415                vec![
2416                    DfRecordBatch::try_new(
2417                        schema.clone(),
2418                        vec![new_ts_array(unit, vec![94, 95, 96, 97, 98, 99])],
2419                    )
2420                    .unwrap(),
2421                ],
2422            ),
2423            (
2424                PartitionRange {
2425                    start: Timestamp::new(50, unit.into()),
2426                    end: Timestamp::new(98, unit.into()),
2427                    num_rows: 3,
2428                    identifier: 1,
2429                },
2430                vec![
2431                    // Values must be < 70 (outside group 1's range) to avoid ambiguity
2432                    DfRecordBatch::try_new(
2433                        schema.clone(),
2434                        vec![new_ts_array(unit, vec![55, 60, 65])],
2435                    )
2436                    .unwrap(),
2437                ],
2438            ),
2439        ];
2440
2441        // With limit=4, we get [99, 98, 97, 96] from group 1
2442        // Threshold is 96, next group's primary_end is 98
2443        // 96 < 98, so threshold check says "could continue"
2444        // But limit is exhausted (0), so we stop anyway
2445        let expected_output = Some(
2446            DfRecordBatch::try_new(
2447                schema.clone(),
2448                vec![new_ts_array(unit, vec![99, 98, 97, 96])],
2449            )
2450            .unwrap(),
2451        );
2452
2453        // Note: We pull 9 rows (both batches) because we need to read batch 2
2454        // to detect the group boundary, even though we stop after outputting group 1.
2455        run_test(
2456            2003,
2457            input_ranged_data,
2458            schema.clone(),
2459            SortOptions {
2460                descending: true,
2461                ..Default::default()
2462            },
2463            Some(4),
2464            expected_output,
2465            Some(9), // Pull both batches to detect boundary
2466        )
2467        .await;
2468    }
2469
2470    /// Test ascending sort with threshold-based early termination.
2471    #[tokio::test]
2472    async fn test_ascending_threshold_early_termination() {
2473        let unit = TimeUnit::Millisecond;
2474        let schema = Arc::new(Schema::new(vec![Field::new(
2475            "ts",
2476            DataType::Timestamp(unit, None),
2477            false,
2478        )]));
2479
2480        // For ascending: primary_end is start, ranges sorted by (start ASC, end ASC)
2481        // Group 1 (start=10) has 6 rows
2482        // Group 2 (start=20) has 3 rows - should NOT be processed because
2483        // threshold (13) < next_primary_end (20)
2484        let input_ranged_data = vec![
2485            (
2486                PartitionRange {
2487                    start: Timestamp::new(10, unit.into()),
2488                    end: Timestamp::new(50, unit.into()),
2489                    num_rows: 6,
2490                    identifier: 0,
2491                },
2492                vec![
2493                    DfRecordBatch::try_new(
2494                        schema.clone(),
2495                        vec![new_ts_array(unit, vec![10, 11, 12, 13, 14, 15])],
2496                    )
2497                    .unwrap(),
2498                ],
2499            ),
2500            (
2501                PartitionRange {
2502                    start: Timestamp::new(20, unit.into()),
2503                    end: Timestamp::new(60, unit.into()),
2504                    num_rows: 3,
2505                    identifier: 1,
2506                },
2507                vec![
2508                    DfRecordBatch::try_new(
2509                        schema.clone(),
2510                        vec![new_ts_array(unit, vec![25, 30, 35])],
2511                    )
2512                    .unwrap(),
2513                ],
2514            ),
2515            // still read this batch to detect group boundary(?)
2516            (
2517                PartitionRange {
2518                    start: Timestamp::new(60, unit.into()),
2519                    end: Timestamp::new(70, unit.into()),
2520                    num_rows: 2,
2521                    identifier: 1,
2522                },
2523                vec![
2524                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![60, 61])])
2525                        .unwrap(),
2526                ],
2527            ),
2528            // after boundary detected, this following one should not be read
2529            (
2530                PartitionRange {
2531                    start: Timestamp::new(61, unit.into()),
2532                    end: Timestamp::new(70, unit.into()),
2533                    num_rows: 2,
2534                    identifier: 1,
2535                },
2536                vec![
2537                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![71, 72])])
2538                        .unwrap(),
2539                ],
2540            ),
2541        ];
2542
2543        // With limit=4, ascending: top 4 (smallest) from group 1 are [10, 11, 12, 13]
2544        // Threshold is 13 (largest in top-k), next group's primary_end is 20
2545        // Since 13 < 20, we stop after group 1 (no value in group 2 can be < 13)
2546        let expected_output = Some(
2547            DfRecordBatch::try_new(
2548                schema.clone(),
2549                vec![new_ts_array(unit, vec![10, 11, 12, 13])],
2550            )
2551            .unwrap(),
2552        );
2553
2554        run_test(
2555            2004,
2556            input_ranged_data,
2557            schema.clone(),
2558            SortOptions {
2559                descending: false,
2560                ..Default::default()
2561            },
2562            Some(4),
2563            expected_output,
2564            Some(11), // Pull first two batches to detect boundary
2565        )
2566        .await;
2567    }
2568
2569    #[tokio::test]
2570    async fn test_ascending_threshold_early_termination_case_two() {
2571        let unit = TimeUnit::Millisecond;
2572        let schema = Arc::new(Schema::new(vec![Field::new(
2573            "ts",
2574            DataType::Timestamp(unit, None),
2575            false,
2576        )]));
2577
2578        // For ascending: primary_end is start, ranges sorted by (start ASC, end ASC)
2579        // Group 1 (start=0) has 4 rows, Group 2 (start=4) has 1 row, Group 3 (start=5) has 4 rows
2580        // After reading all data: [9,10,11,12, 21, 5,6,7,8]
2581        // Sorted ascending: [5,6,7,8, 9,10,11,12, 21]
2582        // With limit=4, output should be smallest 4: [5,6,7,8]
2583        // Algorithm continues reading until start=42 > threshold=8, confirming no smaller values exist
2584        let input_ranged_data = vec![
2585            (
2586                PartitionRange {
2587                    start: Timestamp::new(0, unit.into()),
2588                    end: Timestamp::new(20, unit.into()),
2589                    num_rows: 4,
2590                    identifier: 0,
2591                },
2592                vec![
2593                    DfRecordBatch::try_new(
2594                        schema.clone(),
2595                        vec![new_ts_array(unit, vec![9, 10, 11, 12])],
2596                    )
2597                    .unwrap(),
2598                ],
2599            ),
2600            (
2601                PartitionRange {
2602                    start: Timestamp::new(4, unit.into()),
2603                    end: Timestamp::new(25, unit.into()),
2604                    num_rows: 1,
2605                    identifier: 1,
2606                },
2607                vec![
2608                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![21])])
2609                        .unwrap(),
2610                ],
2611            ),
2612            (
2613                PartitionRange {
2614                    start: Timestamp::new(5, unit.into()),
2615                    end: Timestamp::new(25, unit.into()),
2616                    num_rows: 4,
2617                    identifier: 1,
2618                },
2619                vec![
2620                    DfRecordBatch::try_new(
2621                        schema.clone(),
2622                        vec![new_ts_array(unit, vec![5, 6, 7, 8])],
2623                    )
2624                    .unwrap(),
2625                ],
2626            ),
2627            // This still will be read to detect boundary, but should not contribute to output
2628            (
2629                PartitionRange {
2630                    start: Timestamp::new(42, unit.into()),
2631                    end: Timestamp::new(52, unit.into()),
2632                    num_rows: 2,
2633                    identifier: 1,
2634                },
2635                vec![
2636                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![42, 51])])
2637                        .unwrap(),
2638                ],
2639            ),
2640            // This following one should not be read after boundary detected
2641            (
2642                PartitionRange {
2643                    start: Timestamp::new(48, unit.into()),
2644                    end: Timestamp::new(53, unit.into()),
2645                    num_rows: 2,
2646                    identifier: 1,
2647                },
2648                vec![
2649                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![48, 51])])
2650                        .unwrap(),
2651                ],
2652            ),
2653        ];
2654
2655        // With limit=4, ascending: after processing all ranges, smallest 4 are [5, 6, 7, 8]
2656        // Threshold is 8 (4th smallest value), algorithm reads until start=42 > threshold=8
2657        let expected_output = Some(
2658            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![5, 6, 7, 8])])
2659                .unwrap(),
2660        );
2661
2662        run_test(
2663            2005,
2664            input_ranged_data,
2665            schema.clone(),
2666            SortOptions {
2667                descending: false,
2668                ..Default::default()
2669            },
2670            Some(4),
2671            expected_output,
2672            Some(11), // Read first 4 ranges to confirm threshold boundary
2673        )
2674        .await;
2675    }
2676
2677    /// Test early stop behavior with null values in sort column.
2678    /// Verifies that nulls are handled correctly based on nulls_first option.
2679    #[tokio::test]
2680    async fn test_early_stop_with_nulls() {
2681        let unit = TimeUnit::Millisecond;
2682        let schema = Arc::new(Schema::new(vec![Field::new(
2683            "ts",
2684            DataType::Timestamp(unit, None),
2685            true, // nullable
2686        )]));
2687
2688        // Helper function to create nullable timestamp array
2689        let new_nullable_ts_array = |unit: TimeUnit, arr: Vec<Option<i64>>| -> ArrayRef {
2690            match unit {
2691                TimeUnit::Second => Arc::new(TimestampSecondArray::from(arr)) as ArrayRef,
2692                TimeUnit::Millisecond => Arc::new(TimestampMillisecondArray::from(arr)) as ArrayRef,
2693                TimeUnit::Microsecond => Arc::new(TimestampMicrosecondArray::from(arr)) as ArrayRef,
2694                TimeUnit::Nanosecond => Arc::new(TimestampNanosecondArray::from(arr)) as ArrayRef,
2695            }
2696        };
2697
2698        // Test case 1: nulls_first=true, null values should appear first
2699        // Group 1 (end=100): [null, null, 99, 98, 97] -> with limit=3, top 3 are [null, null, 99]
2700        // Threshold is 99, next group end=90, since 99 >= 90, we should stop early
2701        let input_ranged_data = vec![
2702            (
2703                PartitionRange {
2704                    start: Timestamp::new(70, unit.into()),
2705                    end: Timestamp::new(100, unit.into()),
2706                    num_rows: 5,
2707                    identifier: 0,
2708                },
2709                vec![
2710                    DfRecordBatch::try_new(
2711                        schema.clone(),
2712                        vec![new_nullable_ts_array(
2713                            unit,
2714                            vec![Some(99), Some(98), None, Some(97), None],
2715                        )],
2716                    )
2717                    .unwrap(),
2718                ],
2719            ),
2720            (
2721                PartitionRange {
2722                    start: Timestamp::new(50, unit.into()),
2723                    end: Timestamp::new(90, unit.into()),
2724                    num_rows: 3,
2725                    identifier: 1,
2726                },
2727                vec![
2728                    DfRecordBatch::try_new(
2729                        schema.clone(),
2730                        vec![new_nullable_ts_array(
2731                            unit,
2732                            vec![Some(89), Some(88), Some(87)],
2733                        )],
2734                    )
2735                    .unwrap(),
2736                ],
2737            ),
2738        ];
2739
2740        // With nulls_first=true, nulls sort before all values
2741        // For descending, order is: null, null, 99, 98, 97
2742        // With limit=3, we get: null, null, 99
2743        let expected_output = Some(
2744            DfRecordBatch::try_new(
2745                schema.clone(),
2746                vec![new_nullable_ts_array(unit, vec![None, None, Some(99)])],
2747            )
2748            .unwrap(),
2749        );
2750
2751        run_test(
2752            3000,
2753            input_ranged_data,
2754            schema.clone(),
2755            SortOptions {
2756                descending: true,
2757                nulls_first: true,
2758            },
2759            Some(3),
2760            expected_output,
2761            Some(8), // Must read both batches to detect group boundary
2762        )
2763        .await;
2764
2765        // Test case 2: nulls_last=true, null values should appear last
2766        // Group 1 (end=100): [99, 98, 97, null, null] -> with limit=3, top 3 are [99, 98, 97]
2767        // Threshold is 97, next group end=90, since 97 >= 90, we should stop early
2768        let input_ranged_data = vec![
2769            (
2770                PartitionRange {
2771                    start: Timestamp::new(70, unit.into()),
2772                    end: Timestamp::new(100, unit.into()),
2773                    num_rows: 5,
2774                    identifier: 0,
2775                },
2776                vec![
2777                    DfRecordBatch::try_new(
2778                        schema.clone(),
2779                        vec![new_nullable_ts_array(
2780                            unit,
2781                            vec![Some(99), Some(98), Some(97), None, None],
2782                        )],
2783                    )
2784                    .unwrap(),
2785                ],
2786            ),
2787            (
2788                PartitionRange {
2789                    start: Timestamp::new(50, unit.into()),
2790                    end: Timestamp::new(90, unit.into()),
2791                    num_rows: 3,
2792                    identifier: 1,
2793                },
2794                vec![
2795                    DfRecordBatch::try_new(
2796                        schema.clone(),
2797                        vec![new_nullable_ts_array(
2798                            unit,
2799                            vec![Some(89), Some(88), Some(87)],
2800                        )],
2801                    )
2802                    .unwrap(),
2803                ],
2804            ),
2805        ];
2806
2807        // With nulls_last=false (equivalent to nulls_first=false), values sort before nulls
2808        // For descending, order is: 99, 98, 97, null, null
2809        // With limit=3, we get: 99, 98, 97
2810        let expected_output = Some(
2811            DfRecordBatch::try_new(
2812                schema.clone(),
2813                vec![new_nullable_ts_array(
2814                    unit,
2815                    vec![Some(99), Some(98), Some(97)],
2816                )],
2817            )
2818            .unwrap(),
2819        );
2820
2821        run_test(
2822            3001,
2823            input_ranged_data,
2824            schema.clone(),
2825            SortOptions {
2826                descending: true,
2827                nulls_first: false,
2828            },
2829            Some(3),
2830            expected_output,
2831            Some(8), // Must read both batches to detect group boundary
2832        )
2833        .await;
2834    }
2835
2836    /// Test early stop behavior when there's only one group (no next group).
2837    /// In this case, can_stop_early should return false and we should process all data.
2838    #[tokio::test]
2839    async fn test_early_stop_single_group() {
2840        let unit = TimeUnit::Millisecond;
2841        let schema = Arc::new(Schema::new(vec![Field::new(
2842            "ts",
2843            DataType::Timestamp(unit, None),
2844            false,
2845        )]));
2846
2847        // Only one group (all ranges have the same end), no next group to compare against
2848        let input_ranged_data = vec![
2849            (
2850                PartitionRange {
2851                    start: Timestamp::new(70, unit.into()),
2852                    end: Timestamp::new(100, unit.into()),
2853                    num_rows: 6,
2854                    identifier: 0,
2855                },
2856                vec![
2857                    DfRecordBatch::try_new(
2858                        schema.clone(),
2859                        vec![new_ts_array(unit, vec![94, 95, 96, 97, 98, 99])],
2860                    )
2861                    .unwrap(),
2862                ],
2863            ),
2864            (
2865                PartitionRange {
2866                    start: Timestamp::new(50, unit.into()),
2867                    end: Timestamp::new(100, unit.into()),
2868                    num_rows: 3,
2869                    identifier: 1,
2870                },
2871                vec![
2872                    DfRecordBatch::try_new(
2873                        schema.clone(),
2874                        vec![new_ts_array(unit, vec![85, 86, 87])],
2875                    )
2876                    .unwrap(),
2877                ],
2878            ),
2879        ];
2880
2881        // Even though we have enough data in first range, we must process all
2882        // because there's no next group to compare threshold against
2883        let expected_output = Some(
2884            DfRecordBatch::try_new(
2885                schema.clone(),
2886                vec![new_ts_array(unit, vec![99, 98, 97, 96])],
2887            )
2888            .unwrap(),
2889        );
2890
2891        run_test(
2892            3002,
2893            input_ranged_data,
2894            schema.clone(),
2895            SortOptions {
2896                descending: true,
2897                ..Default::default()
2898            },
2899            Some(4),
2900            expected_output,
2901            Some(9), // Must read all batches since no early stop is possible
2902        )
2903        .await;
2904    }
2905
2906    /// Test early stop behavior when threshold exactly equals next group's boundary.
2907    #[tokio::test]
2908    async fn test_early_stop_exact_boundary_equality() {
2909        let unit = TimeUnit::Millisecond;
2910        let schema = Arc::new(Schema::new(vec![Field::new(
2911            "ts",
2912            DataType::Timestamp(unit, None),
2913            false,
2914        )]));
2915
2916        // Test case 1: Descending sort, threshold == next_group_end
2917        // Group 1 (end=100): data up to 90, threshold = 90, next_group_end = 90
2918        // Since 90 >= 90, we should stop early
2919        let input_ranged_data = vec![
2920            (
2921                PartitionRange {
2922                    start: Timestamp::new(70, unit.into()),
2923                    end: Timestamp::new(100, unit.into()),
2924                    num_rows: 4,
2925                    identifier: 0,
2926                },
2927                vec![
2928                    DfRecordBatch::try_new(
2929                        schema.clone(),
2930                        vec![new_ts_array(unit, vec![92, 91, 90, 89])],
2931                    )
2932                    .unwrap(),
2933                ],
2934            ),
2935            (
2936                PartitionRange {
2937                    start: Timestamp::new(50, unit.into()),
2938                    end: Timestamp::new(90, unit.into()),
2939                    num_rows: 3,
2940                    identifier: 1,
2941                },
2942                vec![
2943                    DfRecordBatch::try_new(
2944                        schema.clone(),
2945                        vec![new_ts_array(unit, vec![88, 87, 86])],
2946                    )
2947                    .unwrap(),
2948                ],
2949            ),
2950        ];
2951
2952        let expected_output = Some(
2953            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![92, 91, 90])])
2954                .unwrap(),
2955        );
2956
2957        run_test(
2958            3003,
2959            input_ranged_data,
2960            schema.clone(),
2961            SortOptions {
2962                descending: true,
2963                ..Default::default()
2964            },
2965            Some(3),
2966            expected_output,
2967            Some(7), // Must read both batches to detect boundary
2968        )
2969        .await;
2970
2971        // Test case 2: Ascending sort, threshold == next_group_start
2972        // Group 1 (start=10): data from 10, threshold = 20, next_group_start = 20
2973        // Since 20 < 20 is false, we should continue
2974        let input_ranged_data = vec![
2975            (
2976                PartitionRange {
2977                    start: Timestamp::new(10, unit.into()),
2978                    end: Timestamp::new(50, unit.into()),
2979                    num_rows: 4,
2980                    identifier: 0,
2981                },
2982                vec![
2983                    DfRecordBatch::try_new(
2984                        schema.clone(),
2985                        vec![new_ts_array(unit, vec![10, 15, 20, 25])],
2986                    )
2987                    .unwrap(),
2988                ],
2989            ),
2990            (
2991                PartitionRange {
2992                    start: Timestamp::new(20, unit.into()),
2993                    end: Timestamp::new(60, unit.into()),
2994                    num_rows: 3,
2995                    identifier: 1,
2996                },
2997                vec![
2998                    DfRecordBatch::try_new(
2999                        schema.clone(),
3000                        vec![new_ts_array(unit, vec![21, 22, 23])],
3001                    )
3002                    .unwrap(),
3003                ],
3004            ),
3005        ];
3006
3007        let expected_output = Some(
3008            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![10, 15, 20])])
3009                .unwrap(),
3010        );
3011
3012        run_test(
3013            3004,
3014            input_ranged_data,
3015            schema.clone(),
3016            SortOptions {
3017                descending: false,
3018                ..Default::default()
3019            },
3020            Some(3),
3021            expected_output,
3022            Some(7), // Must read both batches since 20 is not < 20
3023        )
3024        .await;
3025    }
3026
3027    /// Test early stop behavior with empty partition groups.
3028    #[tokio::test]
3029    async fn test_early_stop_with_empty_partitions() {
3030        let unit = TimeUnit::Millisecond;
3031        let schema = Arc::new(Schema::new(vec![Field::new(
3032            "ts",
3033            DataType::Timestamp(unit, None),
3034            false,
3035        )]));
3036
3037        // Test case 1: First group is empty, second group has data
3038        let input_ranged_data = vec![
3039            (
3040                PartitionRange {
3041                    start: Timestamp::new(70, unit.into()),
3042                    end: Timestamp::new(100, unit.into()),
3043                    num_rows: 0,
3044                    identifier: 0,
3045                },
3046                vec![
3047                    // Empty batch for first range
3048                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![])])
3049                        .unwrap(),
3050                ],
3051            ),
3052            (
3053                PartitionRange {
3054                    start: Timestamp::new(50, unit.into()),
3055                    end: Timestamp::new(100, unit.into()),
3056                    num_rows: 0,
3057                    identifier: 1,
3058                },
3059                vec![
3060                    // Empty batch for second range
3061                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![])])
3062                        .unwrap(),
3063                ],
3064            ),
3065            (
3066                PartitionRange {
3067                    start: Timestamp::new(30, unit.into()),
3068                    end: Timestamp::new(80, unit.into()),
3069                    num_rows: 4,
3070                    identifier: 2,
3071                },
3072                vec![
3073                    DfRecordBatch::try_new(
3074                        schema.clone(),
3075                        vec![new_ts_array(unit, vec![74, 75, 76, 77])],
3076                    )
3077                    .unwrap(),
3078                ],
3079            ),
3080            (
3081                PartitionRange {
3082                    start: Timestamp::new(10, unit.into()),
3083                    end: Timestamp::new(60, unit.into()),
3084                    num_rows: 3,
3085                    identifier: 3,
3086                },
3087                vec![
3088                    DfRecordBatch::try_new(
3089                        schema.clone(),
3090                        vec![new_ts_array(unit, vec![58, 59, 60])],
3091                    )
3092                    .unwrap(),
3093                ],
3094            ),
3095        ];
3096
3097        // Group 1 (end=100) is empty, Group 2 (end=80) has data
3098        // Should continue to Group 2 since Group 1 has no data
3099        let expected_output = Some(
3100            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![77, 76])]).unwrap(),
3101        );
3102
3103        run_test(
3104            3005,
3105            input_ranged_data,
3106            schema.clone(),
3107            SortOptions {
3108                descending: true,
3109                ..Default::default()
3110            },
3111            Some(2),
3112            expected_output,
3113            Some(7), // Must read until finding actual data
3114        )
3115        .await;
3116
3117        // Test case 2: Empty partitions between data groups
3118        let input_ranged_data = vec![
3119            (
3120                PartitionRange {
3121                    start: Timestamp::new(70, unit.into()),
3122                    end: Timestamp::new(100, unit.into()),
3123                    num_rows: 4,
3124                    identifier: 0,
3125                },
3126                vec![
3127                    DfRecordBatch::try_new(
3128                        schema.clone(),
3129                        vec![new_ts_array(unit, vec![96, 97, 98, 99])],
3130                    )
3131                    .unwrap(),
3132                ],
3133            ),
3134            (
3135                PartitionRange {
3136                    start: Timestamp::new(50, unit.into()),
3137                    end: Timestamp::new(90, unit.into()),
3138                    num_rows: 0,
3139                    identifier: 1,
3140                },
3141                vec![
3142                    // Empty range - should be skipped
3143                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![])])
3144                        .unwrap(),
3145                ],
3146            ),
3147            (
3148                PartitionRange {
3149                    start: Timestamp::new(30, unit.into()),
3150                    end: Timestamp::new(70, unit.into()),
3151                    num_rows: 0,
3152                    identifier: 2,
3153                },
3154                vec![
3155                    // Another empty range
3156                    DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![])])
3157                        .unwrap(),
3158                ],
3159            ),
3160            (
3161                PartitionRange {
3162                    start: Timestamp::new(10, unit.into()),
3163                    end: Timestamp::new(50, unit.into()),
3164                    num_rows: 3,
3165                    identifier: 3,
3166                },
3167                vec![
3168                    DfRecordBatch::try_new(
3169                        schema.clone(),
3170                        vec![new_ts_array(unit, vec![48, 49, 50])],
3171                    )
3172                    .unwrap(),
3173                ],
3174            ),
3175        ];
3176
3177        // With limit=2 from group 1: [99, 98], threshold=98, next group end=50
3178        // Since 98 >= 50, we should stop early
3179        let expected_output = Some(
3180            DfRecordBatch::try_new(schema.clone(), vec![new_ts_array(unit, vec![99, 98])]).unwrap(),
3181        );
3182
3183        run_test(
3184            3006,
3185            input_ranged_data,
3186            schema.clone(),
3187            SortOptions {
3188                descending: true,
3189                ..Default::default()
3190            },
3191            Some(2),
3192            expected_output,
3193            Some(7), // Must read to detect early stop condition
3194        )
3195        .await;
3196    }
3197}