1use 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
58fn get_primary_end(range: &PartitionRange, descending: bool) -> Timestamp {
63 if descending { range.end } else { range.start }
64}
65
66fn 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 groups.push((current_primary_end, group_start, idx));
88 group_start = idx;
90 current_primary_end = primary_end;
91 }
92 }
93 groups.push((current_primary_end, group_start, ranges.len()));
95
96 groups
97}
98
99#[derive(Debug, Clone)]
107pub struct PartSortExec {
108 expression: PhysicalSortExpr,
110 limit: Option<usize>,
111 input: Arc<dyn ExecutionPlan>,
112 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 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 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)] 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 range_groups: Vec<(Timestamp, usize, usize)>,
370 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 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 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 fn check_in_range(
472 &self,
473 sort_column: &ArrayRef,
474 min_max_idx: (usize, usize),
475 ) -> datafusion_common::Result<()> {
476 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 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 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 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 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 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 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 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 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 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, identifier: 0, })
860 }
861
862 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 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 self.reservation.shrink(2 * total_mem);
955 Ok(sorted)
956 }
957
958 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 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 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 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 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 self.cur_part_idx += 1;
1058
1059 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 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 self.evaluating_batch = Some(remaining_range);
1083 } else if remaining_range.num_rows() != 0 {
1084 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 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 self.cur_part_idx += 1;
1121
1122 if self.cur_part_idx >= self.partition_ranges.len() {
1124 debug_assert!(remaining_range.num_rows() == 0);
1126
1127 return self.sorted_buffer_if_non_empty();
1129 }
1130
1131 if self.is_in_current_group(self.cur_part_idx) {
1133 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 self.evaluating_batch = Some(remaining_range);
1138 } else {
1139 if remaining_range.num_rows() != 0 {
1141 self.push_buffer(remaining_range, sort_column.data_type())?;
1142 }
1143 }
1144 return Ok(None);
1146 }
1147
1148 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 self.evaluating_batch = Some(remaining_range);
1157 } else {
1158 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 let Some(evaluating_batch) = self.evaluating_batch.take()
1185 && evaluating_batch.num_rows() != 0
1186 {
1187 if self.cur_part_idx >= self.partition_ranges.len() {
1189 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 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 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 let part_cnt_bound = 100;
1265 let range_size_bound = 100;
1267 let range_offset_bound = 100;
1268 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 for part_id in 0..rng.usize(0..part_cnt_bound) {
1309 let (start, end) = if descending {
1311 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 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 continue;
1348 }
1349 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 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 (
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 (
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 #[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 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 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 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 #[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 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 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 #[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 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 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 #[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 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 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 #[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 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 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), )
2389 .await;
2390 }
2391
2392 #[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 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 DfRecordBatch::try_new(
2433 schema.clone(),
2434 vec![new_ts_array(unit, vec![55, 60, 65])],
2435 )
2436 .unwrap(),
2437 ],
2438 ),
2439 ];
2440
2441 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 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), )
2467 .await;
2468 }
2469
2470 #[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 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 (
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 (
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 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), )
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 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 (
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 (
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 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), )
2674 .await;
2675 }
2676
2677 #[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, )]));
2687
2688 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 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 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), )
2763 .await;
2764
2765 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 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), )
2833 .await;
2834 }
2835
2836 #[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 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 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), )
2903 .await;
2904 }
2905
2906 #[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 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), )
2969 .await;
2970
2971 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), )
3024 .await;
3025 }
3026
3027 #[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 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 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 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 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), )
3115 .await;
3116
3117 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 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 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 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), )
3195 .await;
3196 }
3197}