use crate::physical_plan::metrics::BaselineMetrics;
use crate::physical_plan::sorts::builder::BatchBuilder;
use crate::physical_plan::sorts::cursor::Cursor;
use crate::physical_plan::sorts::stream::{
FieldCursorStream, PartitionedStream, RowCursorStream,
};
use crate::physical_plan::{
PhysicalSortExpr, RecordBatchStream, SendableRecordBatchStream,
};
use arrow::datatypes::{DataType, SchemaRef};
use arrow::record_batch::RecordBatch;
use arrow_array::*;
use datafusion_common::Result;
use datafusion_execution::memory_pool::MemoryReservation;
use futures::Stream;
use std::pin::Pin;
use std::task::{ready, Context, Poll};
macro_rules! primitive_merge_helper {
($t:ty, $($v:ident),+) => {
merge_helper!(PrimitiveArray<$t>, $($v),+)
};
}
macro_rules! merge_helper {
($t:ty, $sort:ident, $streams:ident, $schema:ident, $tracking_metrics:ident, $batch_size:ident, $fetch:ident, $reservation:ident) => {{
let streams = FieldCursorStream::<$t>::new($sort, $streams);
return Ok(Box::pin(SortPreservingMergeStream::new(
Box::new(streams),
$schema,
$tracking_metrics,
$batch_size,
$fetch,
$reservation,
)));
}};
}
pub fn streaming_merge(
streams: Vec<SendableRecordBatchStream>,
schema: SchemaRef,
expressions: &[PhysicalSortExpr],
metrics: BaselineMetrics,
batch_size: usize,
fetch: Option<usize>,
reservation: MemoryReservation,
) -> Result<SendableRecordBatchStream> {
if expressions.len() == 1 {
let sort = expressions[0].clone();
let data_type = sort.expr.data_type(schema.as_ref())?;
downcast_primitive! {
data_type => (primitive_merge_helper, sort, streams, schema, metrics, batch_size, fetch, reservation),
DataType::Utf8 => merge_helper!(StringArray, sort, streams, schema, metrics, batch_size, fetch, reservation)
DataType::LargeUtf8 => merge_helper!(LargeStringArray, sort, streams, schema, metrics, batch_size, fetch, reservation)
DataType::Binary => merge_helper!(BinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation)
DataType::LargeBinary => merge_helper!(LargeBinaryArray, sort, streams, schema, metrics, batch_size, fetch, reservation)
_ => {}
}
}
let streams = RowCursorStream::try_new(
schema.as_ref(),
expressions,
streams,
reservation.new_empty(),
)?;
Ok(Box::pin(SortPreservingMergeStream::new(
Box::new(streams),
schema,
metrics,
batch_size,
fetch,
reservation,
)))
}
type CursorStream<C> = Box<dyn PartitionedStream<Output = Result<(C, RecordBatch)>>>;
#[derive(Debug)]
struct SortPreservingMergeStream<C> {
in_progress: BatchBuilder,
streams: CursorStream<C>,
metrics: BaselineMetrics,
aborted: bool,
loser_tree: Vec<usize>,
loser_tree_adjusted: bool,
batch_size: usize,
cursors: Vec<Option<C>>,
fetch: Option<usize>,
produced: usize,
}
impl<C: Cursor> SortPreservingMergeStream<C> {
fn new(
streams: CursorStream<C>,
schema: SchemaRef,
metrics: BaselineMetrics,
batch_size: usize,
fetch: Option<usize>,
reservation: MemoryReservation,
) -> Self {
let stream_count = streams.partitions();
Self {
in_progress: BatchBuilder::new(schema, stream_count, batch_size, reservation),
streams,
metrics,
aborted: false,
cursors: (0..stream_count).map(|_| None).collect(),
loser_tree: vec![],
loser_tree_adjusted: false,
batch_size,
fetch,
produced: 0,
}
}
fn maybe_poll_stream(
&mut self,
cx: &mut Context<'_>,
idx: usize,
) -> Poll<Result<()>> {
if self.cursors[idx].is_some() {
return Poll::Ready(Ok(()));
}
match futures::ready!(self.streams.poll_next(cx, idx)) {
None => Poll::Ready(Ok(())),
Some(Err(e)) => Poll::Ready(Err(e)),
Some(Ok((cursor, batch))) => {
self.cursors[idx] = Some(cursor);
Poll::Ready(self.in_progress.push_batch(idx, batch))
}
}
}
fn poll_next_inner(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Option<Result<RecordBatch>>> {
if self.aborted {
return Poll::Ready(None);
}
if self.loser_tree.is_empty() {
for i in 0..self.streams.partitions() {
if let Err(e) = ready!(self.maybe_poll_stream(cx, i)) {
self.aborted = true;
return Poll::Ready(Some(Err(e)));
}
}
self.init_loser_tree();
}
let elapsed_compute = self.metrics.elapsed_compute().clone();
let _timer = elapsed_compute.timer();
loop {
if !self.loser_tree_adjusted {
let winner = self.loser_tree[0];
if let Err(e) = ready!(self.maybe_poll_stream(cx, winner)) {
self.aborted = true;
return Poll::Ready(Some(Err(e)));
}
self.update_loser_tree();
}
let stream_idx = self.loser_tree[0];
if self.advance(stream_idx) {
self.loser_tree_adjusted = false;
self.in_progress.push_row(stream_idx);
if self.fetch_reached() {
self.aborted = true;
} else if self.in_progress.len() < self.batch_size {
continue;
}
}
self.produced += self.in_progress.len();
return Poll::Ready(self.in_progress.build_record_batch().transpose());
}
}
fn fetch_reached(&mut self) -> bool {
self.fetch
.map(|fetch| self.produced + self.in_progress.len() >= fetch)
.unwrap_or(false)
}
fn advance(&mut self, stream_idx: usize) -> bool {
let slot = &mut self.cursors[stream_idx];
match slot.as_mut() {
Some(c) => {
c.advance();
if c.is_finished() {
*slot = None;
}
true
}
None => false,
}
}
#[inline]
fn is_gt(&self, a: usize, b: usize) -> bool {
match (&self.cursors[a], &self.cursors[b]) {
(None, _) => true,
(_, None) => false,
(Some(ac), Some(bc)) => ac.cmp(bc).then_with(|| a.cmp(&b)).is_gt(),
}
}
#[inline]
fn lt_leaf_node_index(&self, cursor_index: usize) -> usize {
(self.cursors.len() + cursor_index) / 2
}
#[inline]
fn lt_parent_node_index(&self, node_idx: usize) -> usize {
node_idx / 2
}
fn init_loser_tree(&mut self) {
self.loser_tree = vec![usize::MAX; self.cursors.len()];
for i in 0..self.cursors.len() {
let mut winner = i;
let mut cmp_node = self.lt_leaf_node_index(i);
while cmp_node != 0 && self.loser_tree[cmp_node] != usize::MAX {
let challenger = self.loser_tree[cmp_node];
if self.is_gt(winner, challenger) {
self.loser_tree[cmp_node] = winner;
winner = challenger;
}
cmp_node = self.lt_parent_node_index(cmp_node);
}
self.loser_tree[cmp_node] = winner;
}
self.loser_tree_adjusted = true;
}
fn update_loser_tree(&mut self) {
let mut winner = self.loser_tree[0];
let mut cmp_node = self.lt_leaf_node_index(winner);
while cmp_node != 0 {
let challenger = self.loser_tree[cmp_node];
if self.is_gt(winner, challenger) {
self.loser_tree[cmp_node] = winner;
winner = challenger;
}
cmp_node = self.lt_parent_node_index(cmp_node);
}
self.loser_tree[0] = winner;
self.loser_tree_adjusted = true;
}
}
impl<C: Cursor + Unpin> Stream for SortPreservingMergeStream<C> {
type Item = Result<RecordBatch>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
let poll = self.poll_next_inner(cx);
self.metrics.record_poll(poll)
}
}
impl<C: Cursor + Unpin> RecordBatchStream for SortPreservingMergeStream<C> {
fn schema(&self) -> SchemaRef {
self.in_progress.schema().clone()
}
}