Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
241 changes: 77 additions & 164 deletions datafusion/physical-plan/src/joins/cross_join.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand All @@ -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};
Expand All @@ -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)]
Expand Down Expand Up @@ -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 {
Expand All @@ -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,
Expand Down Expand Up @@ -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)?;

Expand All @@ -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<usize>) -> Vec<ChildStats> {
Expand Down Expand Up @@ -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<T> {
struct CrossJoinStream {
/// Input schema
schema: Arc<Schema>,
/// Future for data from left side
left_fut: OnceFut<JoinLeftData>,
/// 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<T: BatchTransformer + Unpin + Send> RecordBatchStream for CrossJoinStream<T> {
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(
Expand Down Expand Up @@ -661,119 +622,71 @@ fn build_batch(
.map_err(Into::into)
}

#[async_trait]
impl<T: BatchTransformer + Unpin + Send> Stream for CrossJoinStream<T> {
type Item = Result<RecordBatch>;

fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> Poll<Option<Self::Item>> {
self.poll_next_impl(cx)
}
}

impl<T: BatchTransformer> CrossJoinStream<T> {
/// 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<Option<Result<RecordBatch>>> {
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<RecordBatch, DataFusionError>,
) -> 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<Result<StatefulStreamResult<Option<RecordBatch>>>> {
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<bool> {
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<Result<StatefulStreamResult<Option<RecordBatch>>>> {
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<Option<RecordBatch>> {
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<StatefulStreamResult<Option<RecordBatch>>> {
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<RecordBatch, DataFusionError>,
) -> 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(())
}
}

Expand Down
Loading