diff --git a/datafusion/physical-plan/src/joins/cross_join.rs b/datafusion/physical-plan/src/joins/cross_join.rs index 8a477c1021d1b..3e798573f5dec 100644 --- a/datafusion/physical-plan/src/joins/cross_join.rs +++ b/datafusion/physical-plan/src/joins/cross_join.rs @@ -18,11 +18,11 @@ //! Defines the cross join plan for loading the left side of the cross join //! and producing batches in parallel for the right partitions -use std::{sync::Arc, task::Poll}; +use std::future::poll_fn; +use std::sync::Arc; use super::utils::{ - BatchSplitter, BatchTransformer, BuildProbeJoinMetrics, NoopBatchTransformer, - OnceAsync, OnceFut, StatefulStreamResult, adjust_right_output_partitioning, + BuildProbeJoinMetrics, OnceAsync, OnceFut, adjust_right_output_partitioning, reorder_output_after_swap, }; use crate::execution_plan::{EmissionType, boundedness_from_children}; @@ -32,12 +32,11 @@ use crate::projection::{ physical_to_column_exprs, }; use crate::statistics::{ChildStats, StatisticsArgs}; -use crate::stream::EmptyRecordBatchStream; +use crate::stream::{EmptyRecordBatchStream, ObservedStream, RecordBatchStreamAdapter}; use crate::{ ChildrenPropertiesMode, ColumnStatistics, DisplayAs, DisplayFormatType, Distribution, - ExecutionPlan, ExecutionPlanProperties, PlanProperties, RecordBatchStream, - ReplaceChildrenOptions, SendableRecordBatchStream, Statistics, handle_state, - validate_child_count, + ExecutionPlan, ExecutionPlanProperties, PlanProperties, ReplaceChildrenOptions, + SendableRecordBatchStream, Statistics, validate_child_count, }; use arrow::array::{RecordBatch, RecordBatchOptions}; @@ -46,15 +45,15 @@ use arrow::datatypes::{Fields, Schema, SchemaRef}; use datafusion_common::stats::Precision; use datafusion_common::tree_node::TreeNodeRecursion; use datafusion_common::{ - JoinType, Result, ScalarValue, assert_eq_or_internal_err, internal_err, + DataFusionError, JoinType, Result, ScalarValue, assert_eq_or_internal_err, }; -use datafusion_execution::TaskContext; use datafusion_execution::memory_pool::{MemoryConsumer, MemoryReservation}; +use datafusion_execution::{TaskContext, TryEmitter, async_try_stream}; use datafusion_physical_expr::PhysicalExpr; use datafusion_physical_expr::equivalence::join_equivalence_properties; -use async_trait::async_trait; -use futures::{Stream, StreamExt, TryStreamExt, ready}; +use futures::{StreamExt, TryStreamExt}; +use num_traits::Zero; /// Data of the left side that is buffered into memory #[derive(Debug)] @@ -208,7 +207,7 @@ async fn load_left_input( let left_schema = stream.schema(); // Load all batches and count the rows - let (batches, _metrics, reservation) = stream + let (batches, metrics, reservation) = stream .try_fold( (Vec::new(), metrics, reservation), |(mut batches, metrics, reservation), batch| async { @@ -226,7 +225,9 @@ async fn load_left_input( ) .await?; + let build_timer = metrics.build_time.timer(); let merged_batch = concat_batches(&left_schema, &batches)?; + build_timer.done(); Ok(JoinLeftData { merged_batch, @@ -367,10 +368,6 @@ impl ExecutionPlan for CrossJoinExec { let reservation = MemoryConsumer::new("CrossJoinExec").register(context.memory_pool()); - let batch_size = context.session_config().batch_size(); - let enforce_batch_size_in_joins = - context.session_config().enforce_batch_size_in_joins(); - let left_fut = self.left_fut.try_once(|| { let left_stream = self.left.execute(0, context)?; @@ -381,29 +378,24 @@ impl ExecutionPlan for CrossJoinExec { )) })?; - if enforce_batch_size_in_joins { - Ok(Box::pin(CrossJoinStream { - schema: Arc::clone(&self.schema), - left_fut, - right: stream, - left_index: 0, - join_metrics, - state: CrossJoinStreamState::WaitBuildSide, - left_data: RecordBatch::new_empty(self.left().schema()), - batch_transformer: BatchSplitter::new(batch_size), - })) - } else { - Ok(Box::pin(CrossJoinStream { - schema: Arc::clone(&self.schema), - left_fut, - right: stream, - left_index: 0, - join_metrics, - state: CrossJoinStreamState::WaitBuildSide, - left_data: RecordBatch::new_empty(self.left().schema()), - batch_transformer: NoopBatchTransformer::new(), - })) - } + let mut state = CrossJoinStream { + schema: Arc::clone(&self.schema), + left_fut, + right: stream, + join_metrics, + left_data: RecordBatch::new_empty(self.left().schema()), + }; + + let schema = Arc::clone(&self.schema); + let baseline_metrics = state.join_metrics.baseline.clone(); + let stream = + async_try_stream(|mut emitter| async move { state.join(&mut emitter).await }); + + Ok(Box::pin(ObservedStream::new( + Box::pin(RecordBatchStreamAdapter::new(schema, stream)), + baseline_metrics, + None, + ))) } fn child_stats_requests(&self, partition: Option) -> Vec { @@ -589,48 +581,17 @@ fn stats_cartesian_product( } /// A stream that issues [RecordBatch]es as they arrive from the right of the join. -struct CrossJoinStream { +struct CrossJoinStream { /// Input schema schema: Arc, /// Future for data from left side left_fut: OnceFut, /// Right side stream right: SendableRecordBatchStream, - /// Current value on the left - left_index: usize, /// Join execution metrics join_metrics: BuildProbeJoinMetrics, - /// State of the stream - state: CrossJoinStreamState, /// Left data (copy of the entire buffered left side) left_data: RecordBatch, - /// Batch transformer - batch_transformer: T, -} - -impl RecordBatchStream for CrossJoinStream { - fn schema(&self) -> SchemaRef { - Arc::clone(&self.schema) - } -} - -/// Represents states of CrossJoinStream -enum CrossJoinStreamState { - WaitBuildSide, - FetchProbeBatch, - /// Holds the currently processed right side batch - BuildBatches(RecordBatch), -} - -impl CrossJoinStreamState { - /// Tries to extract RecordBatch from CrossJoinStreamState enum. - /// Returns an error if state is not BuildBatches state. - fn try_as_record_batch(&mut self) -> Result<&RecordBatch> { - match self { - CrossJoinStreamState::BuildBatches(rb) => Ok(rb), - _ => internal_err!("Expected RecordBatch in BuildBatches state"), - } - } } fn build_batch( @@ -661,119 +622,71 @@ fn build_batch( .map_err(Into::into) } -#[async_trait] -impl Stream for CrossJoinStream { - type Item = Result; - - fn poll_next( - mut self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - ) -> Poll> { - self.poll_next_impl(cx) - } -} - -impl CrossJoinStream { - /// Separate implementation function that unpins the [`CrossJoinStream`] so - /// that partial borrows work correctly - fn poll_next_impl( +impl CrossJoinStream { + // Collect the left (build) side, then continue processing the right side against it until we have no more rows on the right + async fn join( &mut self, - cx: &mut std::task::Context<'_>, - ) -> Poll>> { - loop { - return match self.state { - CrossJoinStreamState::WaitBuildSide => { - handle_state!(ready!(self.collect_build_side(cx))) - } - CrossJoinStreamState::FetchProbeBatch => { - handle_state!(ready!(self.fetch_probe_batch(cx))) - } - CrossJoinStreamState::BuildBatches(_) => { - let poll = handle_state!(self.build_batches()); - self.join_metrics.baseline.record_poll(poll) - } - }; + emitter: &mut TryEmitter, + ) -> Result<()> { + if !self.collect_build_side().await? { + return Ok(()); } + + self.process_right_batch(emitter).await?; + + Ok(()) } - /// Collects build (left) side of the join into the state. In case of an empty build batch, - /// the execution terminates. Otherwise, the state is updated to fetch probe (right) batch. - fn collect_build_side( - &mut self, - cx: &mut std::task::Context<'_>, - ) -> Poll>>> { - let build_timer = self.join_metrics.build_time.timer(); - let left_data = match ready!(self.left_fut.get(cx)) { - Ok(left_data) => left_data, - Err(e) => return Poll::Ready(Err(e)), - }; - build_timer.done(); - - let left_data = left_data.merged_batch.clone(); - let result = if left_data.num_rows() == 0 { - StatefulStreamResult::Ready(None) - } else { - self.left_data = left_data; - self.state = CrossJoinStreamState::FetchProbeBatch; - StatefulStreamResult::Continue - }; - Poll::Ready(Ok(result)) + /// Collects build (left) side of the join into the state. In case of an empty build batch, the execution terminates. + /// Returns true if build side was loaded and non-empty + async fn collect_build_side(&mut self) -> Result { + let left_data = poll_fn(|cx| { + self.left_fut + .get(cx) + .map(|res| res.map(|data| data.merged_batch.clone())) + }) + .await?; + + let is_empty = left_data.num_rows().is_zero(); + self.left_data = left_data; + Ok(!is_empty) } - /// Fetches the probe (right) batch, updates the metrics, and save the batch in the state. - /// Then, the state is updated to build result batches. - fn fetch_probe_batch( - &mut self, - cx: &mut std::task::Context<'_>, - ) -> Poll>>> { - self.left_index = 0; - let right_data = match ready!(self.right.poll_next_unpin(cx)) { + /// Fetches the probe (right) batch, updates the metrics, and returns the batch + async fn fetch_probe_batch(&mut self) -> Result> { + let right_data = match self.right.next().await { Some(Ok(right_data)) => right_data, - Some(Err(e)) => return Poll::Ready(Err(e)), + Some(Err(e)) => return Err(e), None => { // Release the right (probe) input pipeline's resources. let right_schema = self.right.schema(); self.right = Box::pin(EmptyRecordBatchStream::new(right_schema)); - return Poll::Ready(Ok(StatefulStreamResult::Ready(None))); + return Ok(None); } }; self.join_metrics.input_batches.add(1); self.join_metrics.input_rows.add(right_data.num_rows()); - self.state = CrossJoinStreamState::BuildBatches(right_data); - Poll::Ready(Ok(StatefulStreamResult::Continue)) + Ok(Some(right_data)) } - /// Joins the indexed row of left data with the current probe batch. - /// If all the results are produced, the state is set to fetch new probe batch. - fn build_batches(&mut self) -> Result>> { - let right_batch = self.state.try_as_record_batch()?; - if self.left_index < self.left_data.num_rows() { - match self.batch_transformer.next() { - None => { - let join_timer = self.join_metrics.join_time.timer(); - let result = build_batch( - self.left_index, - right_batch, - &self.left_data, - &self.schema, - ); - join_timer.done(); - - self.batch_transformer.set_batch(result?); - } - Some((batch, last)) => { - if last { - self.left_index += 1; - } - - return Ok(StatefulStreamResult::Ready(Some(batch))); - } + /// Joins the left data with the current probe batch, using the emitter to emit the resultant batches + async fn process_right_batch( + &mut self, + emitter: &mut TryEmitter, + ) -> Result<()> { + while let Some(right_batch) = self.fetch_probe_batch().await? { + for left_index in 0..self.left_data.num_rows() { + let join_timer = self.join_metrics.join_time.timer(); + let result = + build_batch(left_index, &right_batch, &self.left_data, &self.schema)?; + join_timer.done(); + + emitter.emit(result).await; } - } else { - self.state = CrossJoinStreamState::FetchProbeBatch; } - Ok(StatefulStreamResult::Continue) + + Ok(()) } }