use std::fmt;
use std::fmt::Debug;
use std::sync::Arc;
use std::task::Poll;
use std::vec;
use std::{any::Any, usize};
use crate::physical_plan::common::SharedMemoryReservation;
use crate::physical_plan::joins::hash_join::{
build_equal_condition_join_indices, update_hash,
};
use crate::physical_plan::joins::hash_join_utils::{
build_filter_expression_graph, calculate_filter_expr_intervals, combine_two_batches,
convert_sort_expr_with_filter_schema, get_pruning_anti_indices,
get_pruning_semi_indices, record_visited_indices, IntervalCalculatorInnerState,
PruningJoinHashMap,
};
use crate::physical_plan::joins::StreamJoinPartitionMode;
use crate::physical_plan::DisplayAs;
use crate::physical_plan::{
expressions::Column,
expressions::PhysicalSortExpr,
joins::{
hash_join_utils::SortedFilterExpr,
utils::{
build_batch_from_indices, build_join_schema, check_join_is_valid,
combine_join_equivalence_properties, partitioned_join_output_partitioning,
ColumnIndex, JoinFilter, JoinOn, JoinSide,
},
},
metrics::{self, ExecutionPlanMetricsSet, MetricBuilder, MetricsSet},
DisplayFormatType, Distribution, EquivalenceProperties, ExecutionPlan, Partitioning,
RecordBatchStream, SendableRecordBatchStream, Statistics,
};
use arrow::array::{ArrowPrimitiveType, NativeAdapter, PrimitiveArray, PrimitiveBuilder};
use arrow::compute::concat_batches;
use arrow::datatypes::{Schema, SchemaRef};
use arrow::record_batch::RecordBatch;
use datafusion_common::utils::bisect;
use datafusion_common::{internal_err, plan_err, JoinType};
use datafusion_common::{DataFusionError, Result};
use datafusion_execution::memory_pool::MemoryConsumer;
use datafusion_execution::TaskContext;
use datafusion_physical_expr::intervals::ExprIntervalGraph;
use ahash::RandomState;
use futures::stream::{select, BoxStream};
use futures::{Stream, StreamExt};
use hashbrown::HashSet;
use parking_lot::Mutex;
const HASHMAP_SHRINK_SCALE_FACTOR: usize = 4;
#[derive(Debug)]
pub struct SymmetricHashJoinExec {
pub(crate) left: Arc<dyn ExecutionPlan>,
pub(crate) right: Arc<dyn ExecutionPlan>,
pub(crate) on: Vec<(Column, Column)>,
pub(crate) filter: Option<JoinFilter>,
pub(crate) join_type: JoinType,
filter_state: Option<Arc<Mutex<IntervalCalculatorInnerState>>>,
schema: SchemaRef,
random_state: RandomState,
metrics: ExecutionPlanMetricsSet,
column_indices: Vec<ColumnIndex>,
pub(crate) null_equals_null: bool,
mode: StreamJoinPartitionMode,
}
#[derive(Debug)]
struct SymmetricHashJoinSideMetrics {
input_batches: metrics::Count,
input_rows: metrics::Count,
}
#[derive(Debug)]
struct SymmetricHashJoinMetrics {
left: SymmetricHashJoinSideMetrics,
right: SymmetricHashJoinSideMetrics,
pub(crate) stream_memory_usage: metrics::Gauge,
output_batches: metrics::Count,
output_rows: metrics::Count,
}
impl SymmetricHashJoinMetrics {
pub fn new(partition: usize, metrics: &ExecutionPlanMetricsSet) -> Self {
let input_batches =
MetricBuilder::new(metrics).counter("input_batches", partition);
let input_rows = MetricBuilder::new(metrics).counter("input_rows", partition);
let left = SymmetricHashJoinSideMetrics {
input_batches,
input_rows,
};
let input_batches =
MetricBuilder::new(metrics).counter("input_batches", partition);
let input_rows = MetricBuilder::new(metrics).counter("input_rows", partition);
let right = SymmetricHashJoinSideMetrics {
input_batches,
input_rows,
};
let stream_memory_usage =
MetricBuilder::new(metrics).gauge("stream_memory_usage", partition);
let output_batches =
MetricBuilder::new(metrics).counter("output_batches", partition);
let output_rows = MetricBuilder::new(metrics).output_rows(partition);
Self {
left,
right,
output_batches,
stream_memory_usage,
output_rows,
}
}
}
impl SymmetricHashJoinExec {
pub fn try_new(
left: Arc<dyn ExecutionPlan>,
right: Arc<dyn ExecutionPlan>,
on: JoinOn,
filter: Option<JoinFilter>,
join_type: &JoinType,
null_equals_null: bool,
mode: StreamJoinPartitionMode,
) -> Result<Self> {
let left_schema = left.schema();
let right_schema = right.schema();
if on.is_empty() {
return plan_err!(
"On constraints in SymmetricHashJoinExec should be non-empty"
);
}
check_join_is_valid(&left_schema, &right_schema, &on)?;
let (schema, column_indices) =
build_join_schema(&left_schema, &right_schema, join_type);
let random_state = RandomState::with_seeds(0, 0, 0, 0);
let filter_state = if filter.is_some() {
let inner_state = IntervalCalculatorInnerState::default();
Some(Arc::new(Mutex::new(inner_state)))
} else {
None
};
Ok(SymmetricHashJoinExec {
left,
right,
on,
filter,
join_type: *join_type,
filter_state,
schema: Arc::new(schema),
random_state,
metrics: ExecutionPlanMetricsSet::new(),
column_indices,
null_equals_null,
mode,
})
}
pub fn left(&self) -> &Arc<dyn ExecutionPlan> {
&self.left
}
pub fn right(&self) -> &Arc<dyn ExecutionPlan> {
&self.right
}
pub fn on(&self) -> &[(Column, Column)] {
&self.on
}
pub fn filter(&self) -> Option<&JoinFilter> {
self.filter.as_ref()
}
pub fn join_type(&self) -> &JoinType {
&self.join_type
}
pub fn null_equals_null(&self) -> bool {
self.null_equals_null
}
pub fn check_if_order_information_available(&self) -> Result<bool> {
if let Some(filter) = self.filter() {
let left = self.left();
if let Some(left_ordering) = left.output_ordering() {
let right = self.right();
if let Some(right_ordering) = right.output_ordering() {
let left_convertible = convert_sort_expr_with_filter_schema(
&JoinSide::Left,
filter,
&left.schema(),
&left_ordering[0],
)?
.is_some();
let right_convertible = convert_sort_expr_with_filter_schema(
&JoinSide::Right,
filter,
&right.schema(),
&right_ordering[0],
)?
.is_some();
return Ok(left_convertible && right_convertible);
}
}
}
Ok(false)
}
}
impl DisplayAs for SymmetricHashJoinExec {
fn fmt_as(&self, t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result {
match t {
DisplayFormatType::Default | DisplayFormatType::Verbose => {
let display_filter = self.filter.as_ref().map_or_else(
|| "".to_string(),
|f| format!(", filter={}", f.expression()),
);
let on = self
.on
.iter()
.map(|(c1, c2)| format!("({}, {})", c1, c2))
.collect::<Vec<String>>()
.join(", ");
write!(
f,
"SymmetricHashJoinExec: mode={:?}, join_type={:?}, on=[{}]{}",
self.mode, self.join_type, on, display_filter
)
}
}
}
}
impl ExecutionPlan for SymmetricHashJoinExec {
fn as_any(&self) -> &dyn Any {
self
}
fn schema(&self) -> SchemaRef {
self.schema.clone()
}
fn unbounded_output(&self, children: &[bool]) -> Result<bool> {
Ok(children.iter().any(|u| *u))
}
fn benefits_from_input_partitioning(&self) -> Vec<bool> {
vec![false, false]
}
fn required_input_distribution(&self) -> Vec<Distribution> {
match self.mode {
StreamJoinPartitionMode::Partitioned => {
let (left_expr, right_expr) = self
.on
.iter()
.map(|(l, r)| (Arc::new(l.clone()) as _, Arc::new(r.clone()) as _))
.unzip();
vec![
Distribution::HashPartitioned(left_expr),
Distribution::HashPartitioned(right_expr),
]
}
StreamJoinPartitionMode::SinglePartition => {
vec![Distribution::SinglePartition, Distribution::SinglePartition]
}
}
}
fn output_partitioning(&self) -> Partitioning {
let left_columns_len = self.left.schema().fields.len();
partitioned_join_output_partitioning(
self.join_type,
self.left.output_partitioning(),
self.right.output_partitioning(),
left_columns_len,
)
}
fn output_ordering(&self) -> Option<&[PhysicalSortExpr]> {
None
}
fn equivalence_properties(&self) -> EquivalenceProperties {
let left_columns_len = self.left.schema().fields.len();
combine_join_equivalence_properties(
self.join_type,
self.left.equivalence_properties(),
self.right.equivalence_properties(),
left_columns_len,
self.on(),
self.schema(),
)
}
fn children(&self) -> Vec<Arc<dyn ExecutionPlan>> {
vec![self.left.clone(), self.right.clone()]
}
fn with_new_children(
self: Arc<Self>,
children: Vec<Arc<dyn ExecutionPlan>>,
) -> Result<Arc<dyn ExecutionPlan>> {
Ok(Arc::new(SymmetricHashJoinExec::try_new(
children[0].clone(),
children[1].clone(),
self.on.clone(),
self.filter.clone(),
&self.join_type,
self.null_equals_null,
self.mode,
)?))
}
fn metrics(&self) -> Option<MetricsSet> {
Some(self.metrics.clone_inner())
}
fn statistics(&self) -> Statistics {
Statistics::default()
}
fn execute(
&self,
partition: usize,
context: Arc<TaskContext>,
) -> Result<SendableRecordBatchStream> {
let left_partitions = self.left.output_partitioning().partition_count();
let right_partitions = self.right.output_partitioning().partition_count();
if left_partitions != right_partitions {
return internal_err!(
"Invalid SymmetricHashJoinExec, partition count mismatch {left_partitions}!={right_partitions},\
consider using RepartitionExec"
);
}
let (left_sorted_filter_expr, right_sorted_filter_expr, graph) =
match (&self.filter_state, &self.filter) {
(Some(interval_state), Some(filter)) => build_filter_expression_graph(
interval_state,
&self.left,
&self.right,
filter,
)?,
(_, _) => (None, None, None),
};
let on_left = self.on.iter().map(|on| on.0.clone()).collect::<Vec<_>>();
let on_right = self.on.iter().map(|on| on.1.clone()).collect::<Vec<_>>();
let left_side_joiner =
OneSideHashJoiner::new(JoinSide::Left, on_left, self.left.schema());
let right_side_joiner =
OneSideHashJoiner::new(JoinSide::Right, on_right, self.right.schema());
let left_stream = self
.left
.execute(partition, context.clone())?
.map(|val| (JoinSide::Left, val));
let right_stream = self
.right
.execute(partition, context.clone())?
.map(|val| (JoinSide::Right, val));
let input_stream = select(left_stream, right_stream).boxed();
let reservation = Arc::new(Mutex::new(
MemoryConsumer::new(format!("SymmetricHashJoinStream[{partition}]"))
.register(context.memory_pool()),
));
if let Some(g) = graph.as_ref() {
reservation.lock().try_grow(g.size())?;
}
Ok(Box::pin(SymmetricHashJoinStream {
input_stream,
schema: self.schema(),
filter: self.filter.clone(),
join_type: self.join_type,
random_state: self.random_state.clone(),
left: left_side_joiner,
right: right_side_joiner,
column_indices: self.column_indices.clone(),
metrics: SymmetricHashJoinMetrics::new(partition, &self.metrics),
graph,
left_sorted_filter_expr,
right_sorted_filter_expr,
null_equals_null: self.null_equals_null,
final_result: false,
reservation,
}))
}
}
struct SymmetricHashJoinStream {
input_stream: BoxStream<'static, (JoinSide, Result<RecordBatch>)>,
schema: Arc<Schema>,
filter: Option<JoinFilter>,
join_type: JoinType,
left: OneSideHashJoiner,
right: OneSideHashJoiner,
column_indices: Vec<ColumnIndex>,
graph: Option<ExprIntervalGraph>,
left_sorted_filter_expr: Option<SortedFilterExpr>,
right_sorted_filter_expr: Option<SortedFilterExpr>,
random_state: RandomState,
null_equals_null: bool,
metrics: SymmetricHashJoinMetrics,
reservation: SharedMemoryReservation,
final_result: bool,
}
impl RecordBatchStream for SymmetricHashJoinStream {
fn schema(&self) -> SchemaRef {
self.schema.clone()
}
}
impl Stream for SymmetricHashJoinStream {
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)
}
}
fn determine_prune_length(
buffer: &RecordBatch,
build_side_filter_expr: &SortedFilterExpr,
) -> Result<usize> {
let origin_sorted_expr = build_side_filter_expr.origin_sorted_expr();
let interval = build_side_filter_expr.interval();
let batch_arr = origin_sorted_expr
.expr
.evaluate(buffer)?
.into_array(buffer.num_rows());
let target = if origin_sorted_expr.options.descending {
interval.upper.value.clone()
} else {
interval.lower.value.clone()
};
bisect::<true>(&[batch_arr], &[target], &[origin_sorted_expr.options])
}
fn need_to_produce_result_in_final(build_side: JoinSide, join_type: JoinType) -> bool {
if build_side == JoinSide::Left {
matches!(
join_type,
JoinType::Left | JoinType::LeftAnti | JoinType::Full | JoinType::LeftSemi
)
} else {
matches!(
join_type,
JoinType::Right | JoinType::RightAnti | JoinType::Full | JoinType::RightSemi
)
}
}
fn calculate_indices_by_join_type<L: ArrowPrimitiveType, R: ArrowPrimitiveType>(
build_side: JoinSide,
prune_length: usize,
visited_rows: &HashSet<usize>,
deleted_offset: usize,
join_type: JoinType,
) -> Result<(PrimitiveArray<L>, PrimitiveArray<R>)>
where
NativeAdapter<L>: From<<L as ArrowPrimitiveType>::Native>,
{
let result = match (build_side, join_type) {
(JoinSide::Left, JoinType::Left | JoinType::LeftAnti)
| (JoinSide::Right, JoinType::Right | JoinType::RightAnti)
| (_, JoinType::Full) => {
let build_unmatched_indices =
get_pruning_anti_indices(prune_length, deleted_offset, visited_rows);
let mut builder =
PrimitiveBuilder::<R>::with_capacity(build_unmatched_indices.len());
builder.append_nulls(build_unmatched_indices.len());
let probe_indices = builder.finish();
(build_unmatched_indices, probe_indices)
}
(JoinSide::Left, JoinType::LeftSemi) | (JoinSide::Right, JoinType::RightSemi) => {
let build_unmatched_indices =
get_pruning_semi_indices(prune_length, deleted_offset, visited_rows);
let mut builder =
PrimitiveBuilder::<R>::with_capacity(build_unmatched_indices.len());
builder.append_nulls(build_unmatched_indices.len());
let probe_indices = builder.finish();
(build_unmatched_indices, probe_indices)
}
_ => unreachable!(),
};
Ok(result)
}
pub(crate) fn build_side_determined_results(
build_hash_joiner: &OneSideHashJoiner,
output_schema: &SchemaRef,
prune_length: usize,
probe_schema: SchemaRef,
join_type: JoinType,
column_indices: &[ColumnIndex],
) -> Result<Option<RecordBatch>> {
if need_to_produce_result_in_final(build_hash_joiner.build_side, join_type) {
let (build_indices, probe_indices) = calculate_indices_by_join_type(
build_hash_joiner.build_side,
prune_length,
&build_hash_joiner.visited_rows,
build_hash_joiner.deleted_offset,
join_type,
)?;
let empty_probe_batch = RecordBatch::new_empty(probe_schema);
build_batch_from_indices(
output_schema.as_ref(),
&build_hash_joiner.input_buffer,
&empty_probe_batch,
&build_indices,
&probe_indices,
column_indices,
build_hash_joiner.build_side,
)
.map(|batch| (batch.num_rows() > 0).then_some(batch))
} else {
Ok(None)
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn join_with_probe_batch(
build_hash_joiner: &mut OneSideHashJoiner,
probe_hash_joiner: &mut OneSideHashJoiner,
schema: &SchemaRef,
join_type: JoinType,
filter: Option<&JoinFilter>,
probe_batch: &RecordBatch,
column_indices: &[ColumnIndex],
random_state: &RandomState,
null_equals_null: bool,
) -> Result<Option<RecordBatch>> {
if build_hash_joiner.input_buffer.num_rows() == 0 || probe_batch.num_rows() == 0 {
return Ok(None);
}
let (build_indices, probe_indices) = build_equal_condition_join_indices(
&build_hash_joiner.hashmap,
&build_hash_joiner.input_buffer,
probe_batch,
&build_hash_joiner.on,
&probe_hash_joiner.on,
random_state,
null_equals_null,
&mut build_hash_joiner.hashes_buffer,
filter,
build_hash_joiner.build_side,
Some(build_hash_joiner.deleted_offset),
)?;
if need_to_produce_result_in_final(build_hash_joiner.build_side, join_type) {
record_visited_indices(
&mut build_hash_joiner.visited_rows,
build_hash_joiner.deleted_offset,
&build_indices,
);
}
if need_to_produce_result_in_final(build_hash_joiner.build_side.negate(), join_type) {
record_visited_indices(
&mut probe_hash_joiner.visited_rows,
probe_hash_joiner.offset,
&probe_indices,
);
}
if matches!(
join_type,
JoinType::LeftAnti
| JoinType::RightAnti
| JoinType::LeftSemi
| JoinType::RightSemi
) {
Ok(None)
} else {
build_batch_from_indices(
schema,
&build_hash_joiner.input_buffer,
probe_batch,
&build_indices,
&probe_indices,
column_indices,
build_hash_joiner.build_side,
)
.map(|batch| (batch.num_rows() > 0).then_some(batch))
}
}
pub struct OneSideHashJoiner {
build_side: JoinSide,
pub input_buffer: RecordBatch,
pub(crate) on: Vec<Column>,
pub(crate) hashmap: PruningJoinHashMap,
pub(crate) hashes_buffer: Vec<u64>,
pub(crate) visited_rows: HashSet<usize>,
pub(crate) offset: usize,
pub(crate) deleted_offset: usize,
}
impl OneSideHashJoiner {
pub fn size(&self) -> usize {
let mut size = 0;
size += std::mem::size_of_val(self);
size += std::mem::size_of_val(&self.build_side);
size += self.input_buffer.get_array_memory_size();
size += std::mem::size_of_val(&self.on);
size += self.hashmap.size();
size += self.hashes_buffer.capacity() * std::mem::size_of::<u64>();
size += self.visited_rows.capacity() * std::mem::size_of::<usize>();
size += std::mem::size_of_val(&self.offset);
size += std::mem::size_of_val(&self.deleted_offset);
size
}
pub fn new(build_side: JoinSide, on: Vec<Column>, schema: SchemaRef) -> Self {
Self {
build_side,
input_buffer: RecordBatch::new_empty(schema),
on,
hashmap: PruningJoinHashMap::with_capacity(0),
hashes_buffer: vec![],
visited_rows: HashSet::new(),
offset: 0,
deleted_offset: 0,
}
}
pub(crate) fn update_internal_state(
&mut self,
batch: &RecordBatch,
random_state: &RandomState,
) -> Result<()> {
self.input_buffer = concat_batches(&batch.schema(), [&self.input_buffer, batch])?;
self.hashes_buffer.resize(batch.num_rows(), 0);
update_hash(
&self.on,
batch,
&mut self.hashmap,
self.offset,
random_state,
&mut self.hashes_buffer,
self.deleted_offset,
)?;
Ok(())
}
pub(crate) fn calculate_prune_length_with_probe_batch(
&mut self,
build_side_sorted_filter_expr: &mut SortedFilterExpr,
probe_side_sorted_filter_expr: &mut SortedFilterExpr,
graph: &mut ExprIntervalGraph,
) -> Result<usize> {
if self.input_buffer.num_rows() == 0 {
return Ok(0);
}
let mut filter_intervals = vec![];
for expr in [
&build_side_sorted_filter_expr,
&probe_side_sorted_filter_expr,
] {
filter_intervals.push((expr.node_index(), expr.interval().clone()))
}
graph.update_ranges(&mut filter_intervals)?;
let calculated_build_side_interval = filter_intervals.remove(0).1;
if calculated_build_side_interval.eq(build_side_sorted_filter_expr.interval()) {
return Ok(0);
}
build_side_sorted_filter_expr.set_interval(calculated_build_side_interval);
determine_prune_length(&self.input_buffer, build_side_sorted_filter_expr)
}
pub(crate) fn prune_internal_state(&mut self, prune_length: usize) -> Result<()> {
self.hashmap.prune_hash_values(
prune_length,
self.deleted_offset as u64,
HASHMAP_SHRINK_SCALE_FACTOR,
)?;
for row in self.deleted_offset..(self.deleted_offset + prune_length) {
self.visited_rows.remove(&row);
}
self.input_buffer = self
.input_buffer
.slice(prune_length, self.input_buffer.num_rows() - prune_length);
self.deleted_offset += prune_length;
Ok(())
}
}
impl SymmetricHashJoinStream {
fn size(&self) -> usize {
let mut size = 0;
size += std::mem::size_of_val(&self.input_stream);
size += std::mem::size_of_val(&self.schema);
size += std::mem::size_of_val(&self.filter);
size += std::mem::size_of_val(&self.join_type);
size += self.left.size();
size += self.right.size();
size += std::mem::size_of_val(&self.column_indices);
size += self.graph.as_ref().map(|g| g.size()).unwrap_or(0);
size += std::mem::size_of_val(&self.left_sorted_filter_expr);
size += std::mem::size_of_val(&self.right_sorted_filter_expr);
size += std::mem::size_of_val(&self.random_state);
size += std::mem::size_of_val(&self.null_equals_null);
size += std::mem::size_of_val(&self.metrics);
size += std::mem::size_of_val(&self.final_result);
size
}
fn poll_next_impl(
&mut self,
cx: &mut std::task::Context<'_>,
) -> Poll<Option<Result<RecordBatch>>> {
loop {
match self.input_stream.poll_next_unpin(cx) {
Poll::Ready(Some((side, Ok(probe_batch)))) => {
let (
probe_hash_joiner,
build_hash_joiner,
probe_side_sorted_filter_expr,
build_side_sorted_filter_expr,
probe_side_metrics,
) = if side.eq(&JoinSide::Left) {
(
&mut self.left,
&mut self.right,
&mut self.left_sorted_filter_expr,
&mut self.right_sorted_filter_expr,
&mut self.metrics.left,
)
} else {
(
&mut self.right,
&mut self.left,
&mut self.right_sorted_filter_expr,
&mut self.left_sorted_filter_expr,
&mut self.metrics.right,
)
};
probe_side_metrics.input_batches.add(1);
probe_side_metrics.input_rows.add(probe_batch.num_rows());
probe_hash_joiner
.update_internal_state(&probe_batch, &self.random_state)?;
let equal_result = join_with_probe_batch(
build_hash_joiner,
probe_hash_joiner,
&self.schema,
self.join_type,
self.filter.as_ref(),
&probe_batch,
&self.column_indices,
&self.random_state,
self.null_equals_null,
)?;
probe_hash_joiner.offset += probe_batch.num_rows();
let anti_result = if let (
Some(build_side_sorted_filter_expr),
Some(probe_side_sorted_filter_expr),
Some(graph),
) = (
build_side_sorted_filter_expr.as_mut(),
probe_side_sorted_filter_expr.as_mut(),
self.graph.as_mut(),
) {
calculate_filter_expr_intervals(
&build_hash_joiner.input_buffer,
build_side_sorted_filter_expr,
&probe_batch,
probe_side_sorted_filter_expr,
)?;
let prune_length = build_hash_joiner
.calculate_prune_length_with_probe_batch(
build_side_sorted_filter_expr,
probe_side_sorted_filter_expr,
graph,
)?;
if prune_length > 0 {
let res = build_side_determined_results(
build_hash_joiner,
&self.schema,
prune_length,
probe_batch.schema(),
self.join_type,
&self.column_indices,
)?;
build_hash_joiner.prune_internal_state(prune_length)?;
res
} else {
None
}
} else {
None
};
let result =
combine_two_batches(&self.schema, equal_result, anti_result)?;
let capacity = self.size();
self.metrics.stream_memory_usage.set(capacity);
self.reservation.lock().try_resize(capacity)?;
if let Some(batch) = &result {
self.metrics.output_batches.add(1);
self.metrics.output_rows.add(batch.num_rows());
return Poll::Ready(Ok(result).transpose());
}
}
Poll::Ready(Some((_, Err(e)))) => return Poll::Ready(Some(Err(e))),
Poll::Ready(None) => {
if self.final_result {
return Poll::Ready(None);
}
self.final_result = true;
let left_result = build_side_determined_results(
&self.left,
&self.schema,
self.left.input_buffer.num_rows(),
self.right.input_buffer.schema(),
self.join_type,
&self.column_indices,
)?;
let right_result = build_side_determined_results(
&self.right,
&self.schema,
self.right.input_buffer.num_rows(),
self.left.input_buffer.schema(),
self.join_type,
&self.column_indices,
)?;
let result =
combine_two_batches(&self.schema, left_result, right_result)?;
if let Some(batch) = &result {
self.metrics.output_batches.add(1);
self.metrics.output_rows.add(batch.num_rows());
return Poll::Ready(Ok(result).transpose());
}
}
Poll::Pending => return Poll::Pending,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::compute::SortOptions;
use arrow::datatypes::{DataType, Field, IntervalUnit, Schema, TimeUnit};
use datafusion_execution::config::SessionConfig;
use rstest::*;
use datafusion_expr::Operator;
use datafusion_physical_expr::expressions::{binary, col, Column};
use datafusion_physical_expr::intervals::test_utils::gen_conjunctive_numerical_expr;
use crate::physical_plan::joins::hash_join_utils::tests::complicated_filter;
use crate::physical_plan::joins::test_utils::{
build_sides_record_batches, compare_batches, create_memory_table,
join_expr_tests_fixture_f64, join_expr_tests_fixture_i32,
join_expr_tests_fixture_temporal, partitioned_hash_join_with_filter,
partitioned_sym_join_with_filter,
};
use datafusion_common::ScalarValue;
const TABLE_SIZE: i32 = 1000;
pub async fn experiment(
left: Arc<dyn ExecutionPlan>,
right: Arc<dyn ExecutionPlan>,
filter: Option<JoinFilter>,
join_type: JoinType,
on: JoinOn,
task_ctx: Arc<TaskContext>,
) -> Result<()> {
let first_batches = partitioned_sym_join_with_filter(
left.clone(),
right.clone(),
on.clone(),
filter.clone(),
&join_type,
false,
task_ctx.clone(),
)
.await?;
let second_batches = partitioned_hash_join_with_filter(
left, right, on, filter, &join_type, false, task_ctx,
)
.await?;
compare_batches(&first_batches, &second_batches);
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn complex_join_all_one_ascending_numeric(
#[values(
JoinType::Inner,
JoinType::Left,
JoinType::Right,
JoinType::RightSemi,
JoinType::LeftSemi,
JoinType::LeftAnti,
JoinType::RightAnti,
JoinType::Full
)]
join_type: JoinType,
#[values(
(4, 5),
(11, 21),
(31, 71),
(99, 12),
)]
cardinality: (i32, i32),
) -> Result<()> {
let task_ctx = Arc::new(TaskContext::default());
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: binary(
col("la1", left_schema)?,
Operator::Plus,
col("la2", left_schema)?,
left_schema,
)?,
options: SortOptions::default(),
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("ra1", right_schema)?,
options: SortOptions::default(),
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("0", DataType::Int32, true),
Field::new("1", DataType::Int32, true),
Field::new("2", DataType::Int32, true),
]);
let filter_expr = complicated_filter(&intermediate_schema)?;
let column_indices = vec![
ColumnIndex {
index: 0,
side: JoinSide::Left,
},
ColumnIndex {
index: 4,
side: JoinSide::Left,
},
ColumnIndex {
index: 0,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn join_all_one_ascending_numeric(
#[values(
JoinType::Inner,
JoinType::Left,
JoinType::Right,
JoinType::RightSemi,
JoinType::LeftSemi,
JoinType::LeftAnti,
JoinType::RightAnti,
JoinType::Full
)]
join_type: JoinType,
#[values(
(4, 5),
(11, 21),
(31, 71),
(99, 12),
)]
cardinality: (i32, i32),
#[values(0, 1, 2, 3, 4, 5, 6, 7)] case_expr: usize,
) -> Result<()> {
let task_ctx = Arc::new(TaskContext::default());
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: col("la1", left_schema)?,
options: SortOptions::default(),
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("ra1", right_schema)?,
options: SortOptions::default(),
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Int32, true),
Field::new("right", DataType::Int32, true),
]);
let filter_expr = join_expr_tests_fixture_i32(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
);
let column_indices = vec![
ColumnIndex {
index: 0,
side: JoinSide::Left,
},
ColumnIndex {
index: 0,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn join_all_one_ascending_numeric_v2() -> Result<()> {
let join_type = JoinType::Inner;
let cardinality = (4, 5);
let case_expr = 2;
let task_ctx = Arc::new(TaskContext::default());
let (left_batch, right_batch) = build_sides_record_batches(1000, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: col("la1", left_schema)?,
options: SortOptions::default(),
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("ra1", right_schema)?,
options: SortOptions::default(),
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Int32, true),
Field::new("right", DataType::Int32, true),
]);
let filter_expr = join_expr_tests_fixture_i32(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
);
let column_indices = vec![
ColumnIndex {
index: 0,
side: JoinSide::Left,
},
ColumnIndex {
index: 0,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn join_without_sort_information(
#[values(
JoinType::Inner,
JoinType::Left,
JoinType::Right,
JoinType::RightSemi,
JoinType::LeftSemi,
JoinType::LeftAnti,
JoinType::RightAnti,
JoinType::Full
)]
join_type: JoinType,
#[values(
(4, 5),
(11, 21),
(31, 71),
(99, 12),
)]
cardinality: (i32, i32),
#[values(0, 1, 2, 3, 4, 5, 6)] case_expr: usize,
) -> Result<()> {
let task_ctx = Arc::new(TaskContext::default());
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let (left, right) =
create_memory_table(left_batch, right_batch, vec![], vec![], 13)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Int32, true),
Field::new("right", DataType::Int32, true),
]);
let filter_expr = join_expr_tests_fixture_i32(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
);
let column_indices = vec![
ColumnIndex {
index: 5,
side: JoinSide::Left,
},
ColumnIndex {
index: 5,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn join_without_filter(
#[values(
JoinType::Inner,
JoinType::Left,
JoinType::Right,
JoinType::RightSemi,
JoinType::LeftSemi,
JoinType::LeftAnti,
JoinType::RightAnti,
JoinType::Full
)]
join_type: JoinType,
) -> Result<()> {
let task_ctx = Arc::new(TaskContext::default());
let (left_batch, right_batch) = build_sides_record_batches(TABLE_SIZE, (11, 21))?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let (left, right) =
create_memory_table(left_batch, right_batch, vec![], vec![], 13)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
experiment(left, right, None, join_type, on, task_ctx).await?;
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn join_all_one_descending_numeric_particular(
#[values(
JoinType::Inner,
JoinType::Left,
JoinType::Right,
JoinType::RightSemi,
JoinType::LeftSemi,
JoinType::LeftAnti,
JoinType::RightAnti,
JoinType::Full
)]
join_type: JoinType,
#[values(
(4, 5),
(11, 21),
(31, 71),
(99, 12),
)]
cardinality: (i32, i32),
#[values(0, 1, 2, 3, 4, 5, 6)] case_expr: usize,
) -> Result<()> {
let task_ctx = Arc::new(TaskContext::default());
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: col("la1_des", left_schema)?,
options: SortOptions {
descending: true,
nulls_first: true,
},
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("ra1_des", right_schema)?,
options: SortOptions {
descending: true,
nulls_first: true,
},
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Int32, true),
Field::new("right", DataType::Int32, true),
]);
let filter_expr = join_expr_tests_fixture_i32(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
);
let column_indices = vec![
ColumnIndex {
index: 5,
side: JoinSide::Left,
},
ColumnIndex {
index: 5,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn build_null_columns_first() -> Result<()> {
let join_type = JoinType::Full;
let cardinality = (10, 11);
let case_expr = 1;
let session_config = SessionConfig::new().with_repartition_joins(false);
let task_ctx = TaskContext::default().with_session_config(session_config);
let task_ctx = Arc::new(task_ctx);
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: col("l_asc_null_first", left_schema)?,
options: SortOptions {
descending: false,
nulls_first: true,
},
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("r_asc_null_first", right_schema)?,
options: SortOptions {
descending: false,
nulls_first: true,
},
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Int32, true),
Field::new("right", DataType::Int32, true),
]);
let filter_expr = join_expr_tests_fixture_i32(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
);
let column_indices = vec![
ColumnIndex {
index: 6,
side: JoinSide::Left,
},
ColumnIndex {
index: 6,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn build_null_columns_last() -> Result<()> {
let join_type = JoinType::Full;
let cardinality = (10, 11);
let case_expr = 1;
let session_config = SessionConfig::new().with_repartition_joins(false);
let task_ctx = TaskContext::default().with_session_config(session_config);
let task_ctx = Arc::new(task_ctx);
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: col("l_asc_null_last", left_schema)?,
options: SortOptions {
descending: false,
nulls_first: false,
},
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("r_asc_null_last", right_schema)?,
options: SortOptions {
descending: false,
nulls_first: false,
},
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Int32, true),
Field::new("right", DataType::Int32, true),
]);
let filter_expr = join_expr_tests_fixture_i32(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
);
let column_indices = vec![
ColumnIndex {
index: 7,
side: JoinSide::Left,
},
ColumnIndex {
index: 7,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn build_null_columns_first_descending() -> Result<()> {
let join_type = JoinType::Full;
let cardinality = (10, 11);
let case_expr = 1;
let session_config = SessionConfig::new().with_repartition_joins(false);
let task_ctx = TaskContext::default().with_session_config(session_config);
let task_ctx = Arc::new(task_ctx);
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: col("l_desc_null_first", left_schema)?,
options: SortOptions {
descending: true,
nulls_first: true,
},
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("r_desc_null_first", right_schema)?,
options: SortOptions {
descending: true,
nulls_first: true,
},
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Int32, true),
Field::new("right", DataType::Int32, true),
]);
let filter_expr = join_expr_tests_fixture_i32(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
);
let column_indices = vec![
ColumnIndex {
index: 8,
side: JoinSide::Left,
},
ColumnIndex {
index: 8,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn complex_join_all_one_ascending_numeric_missing_stat() -> Result<()> {
let cardinality = (3, 4);
let join_type = JoinType::Full;
let session_config = SessionConfig::new().with_repartition_joins(false);
let task_ctx = TaskContext::default().with_session_config(session_config);
let task_ctx = Arc::new(task_ctx);
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: col("la1", left_schema)?,
options: SortOptions::default(),
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("ra1", right_schema)?,
options: SortOptions::default(),
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("0", DataType::Int32, true),
Field::new("1", DataType::Int32, true),
Field::new("2", DataType::Int32, true),
]);
let filter_expr = complicated_filter(&intermediate_schema)?;
let column_indices = vec![
ColumnIndex {
index: 0,
side: JoinSide::Left,
},
ColumnIndex {
index: 4,
side: JoinSide::Left,
},
ColumnIndex {
index: 0,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn test_one_side_hash_joiner_visited_rows(
#[values(
(JoinType::Inner, true),
(JoinType::Left,false),
(JoinType::Right, true),
(JoinType::RightSemi, true),
(JoinType::LeftSemi, false),
(JoinType::LeftAnti, false),
(JoinType::RightAnti, true),
(JoinType::Full, false),
)]
case: (JoinType, bool),
) -> Result<()> {
let join_type = case.0;
let should_be_empty = case.1;
let random_state = RandomState::with_seeds(0, 0, 0, 0);
let session_config = SessionConfig::new().with_repartition_joins(false);
let task_ctx = TaskContext::default().with_session_config(session_config);
let task_ctx = Arc::new(task_ctx);
let (left_batch, right_batch) = build_sides_record_batches(20, (1, 1))?;
let left_schema = left_batch.schema();
let right_schema = right_batch.schema();
let (schema, join_column_indices) =
build_join_schema(&left_schema, &right_schema, &join_type);
let join_schema = Arc::new(schema);
let left_sorted = vec![PhysicalSortExpr {
expr: col("la1", &left_schema)?,
options: SortOptions::default(),
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("ra1", &right_schema)?,
options: SortOptions::default(),
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
10,
)?;
let intermediate_schema = Schema::new(vec![
Field::new("0", DataType::Int32, true),
Field::new("1", DataType::Int32, true),
]);
let filter_expr = gen_conjunctive_numerical_expr(
col("0", &intermediate_schema)?,
col("1", &intermediate_schema)?,
(
Operator::Plus,
Operator::Minus,
Operator::Plus,
Operator::Plus,
),
ScalarValue::Int32(Some(0)),
ScalarValue::Int32(Some(3)),
ScalarValue::Int32(Some(0)),
ScalarValue::Int32(Some(3)),
(Operator::Gt, Operator::Lt),
);
let column_indices = vec![
ColumnIndex {
index: 0,
side: JoinSide::Left,
},
ColumnIndex {
index: 0,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
let mut left_side_joiner = OneSideHashJoiner::new(
JoinSide::Left,
vec![Column::new_with_schema("lc1", &left_schema)?],
left_schema,
);
let mut right_side_joiner = OneSideHashJoiner::new(
JoinSide::Right,
vec![Column::new_with_schema("rc1", &right_schema)?],
right_schema,
);
let mut left_stream = left.execute(0, task_ctx.clone())?;
let mut right_stream = right.execute(0, task_ctx)?;
let initial_left_batch = left_stream.next().await.unwrap()?;
left_side_joiner.update_internal_state(&initial_left_batch, &random_state)?;
assert_eq!(
left_side_joiner.input_buffer.num_rows(),
initial_left_batch.num_rows()
);
let initial_right_batch = right_stream.next().await.unwrap()?;
right_side_joiner.update_internal_state(&initial_right_batch, &random_state)?;
assert_eq!(
right_side_joiner.input_buffer.num_rows(),
initial_right_batch.num_rows()
);
join_with_probe_batch(
&mut left_side_joiner,
&mut right_side_joiner,
&join_schema,
join_type,
Some(&filter),
&initial_right_batch,
&join_column_indices,
&random_state,
false,
)?;
assert_eq!(left_side_joiner.visited_rows.is_empty(), should_be_empty);
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn testing_with_temporal_columns(
#[values(
JoinType::Inner,
JoinType::Left,
JoinType::Right,
JoinType::RightSemi,
JoinType::LeftSemi,
JoinType::LeftAnti,
JoinType::RightAnti,
JoinType::Full
)]
join_type: JoinType,
#[values(
(4, 5),
(99, 12),
)]
cardinality: (i32, i32),
#[values(0, 1)] case_expr: usize,
) -> Result<()> {
let session_config = SessionConfig::new().with_repartition_joins(false);
let task_ctx = TaskContext::default().with_session_config(session_config);
let task_ctx = Arc::new(task_ctx);
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let left_sorted = vec![PhysicalSortExpr {
expr: col("lt1", left_schema)?,
options: SortOptions {
descending: false,
nulls_first: true,
},
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("rt1", right_schema)?,
options: SortOptions {
descending: false,
nulls_first: true,
},
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let intermediate_schema = Schema::new(vec![
Field::new(
"left",
DataType::Timestamp(TimeUnit::Millisecond, None),
false,
),
Field::new(
"right",
DataType::Timestamp(TimeUnit::Millisecond, None),
false,
),
]);
let filter_expr = join_expr_tests_fixture_temporal(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
&intermediate_schema,
)?;
let column_indices = vec![
ColumnIndex {
index: 3,
side: JoinSide::Left,
},
ColumnIndex {
index: 3,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn test_with_interval_columns(
#[values(
JoinType::Inner,
JoinType::Left,
JoinType::Right,
JoinType::RightSemi,
JoinType::LeftSemi,
JoinType::LeftAnti,
JoinType::RightAnti,
JoinType::Full
)]
join_type: JoinType,
#[values(
(4, 5),
(99, 12),
)]
cardinality: (i32, i32),
) -> Result<()> {
let session_config = SessionConfig::new().with_repartition_joins(false);
let task_ctx = TaskContext::default().with_session_config(session_config);
let task_ctx = Arc::new(task_ctx);
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let left_sorted = vec![PhysicalSortExpr {
expr: col("li1", left_schema)?,
options: SortOptions {
descending: false,
nulls_first: true,
},
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("ri1", right_schema)?,
options: SortOptions {
descending: false,
nulls_first: true,
},
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Interval(IntervalUnit::DayTime), false),
Field::new("right", DataType::Interval(IntervalUnit::DayTime), false),
]);
let filter_expr = join_expr_tests_fixture_temporal(
0,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
&intermediate_schema,
)?;
let column_indices = vec![
ColumnIndex {
index: 9,
side: JoinSide::Left,
},
ColumnIndex {
index: 9,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn testing_ascending_float_pruning(
#[values(
JoinType::Inner,
JoinType::Left,
JoinType::Right,
JoinType::RightSemi,
JoinType::LeftSemi,
JoinType::LeftAnti,
JoinType::RightAnti,
JoinType::Full
)]
join_type: JoinType,
#[values(
(4, 5),
(99, 12),
)]
cardinality: (i32, i32),
#[values(0, 1, 2, 3, 4, 5, 6, 7)] case_expr: usize,
) -> Result<()> {
let session_config = SessionConfig::new().with_repartition_joins(false);
let task_ctx = TaskContext::default().with_session_config(session_config);
let task_ctx = Arc::new(task_ctx);
let (left_batch, right_batch) =
build_sides_record_batches(TABLE_SIZE, cardinality)?;
let left_schema = &left_batch.schema();
let right_schema = &right_batch.schema();
let left_sorted = vec![PhysicalSortExpr {
expr: col("l_float", left_schema)?,
options: SortOptions::default(),
}];
let right_sorted = vec![PhysicalSortExpr {
expr: col("r_float", right_schema)?,
options: SortOptions::default(),
}];
let (left, right) = create_memory_table(
left_batch,
right_batch,
vec![left_sorted],
vec![right_sorted],
13,
)?;
let on = vec![(
Column::new_with_schema("lc1", left_schema)?,
Column::new_with_schema("rc1", right_schema)?,
)];
let intermediate_schema = Schema::new(vec![
Field::new("left", DataType::Float64, true),
Field::new("right", DataType::Float64, true),
]);
let filter_expr = join_expr_tests_fixture_f64(
case_expr,
col("left", &intermediate_schema)?,
col("right", &intermediate_schema)?,
);
let column_indices = vec![
ColumnIndex {
index: 10, side: JoinSide::Left,
},
ColumnIndex {
index: 10, side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
experiment(left, right, Some(filter), join_type, on, task_ctx).await?;
Ok(())
}
}