use std::collections::{HashMap, VecDeque};
use std::fmt::{Debug, Formatter};
use std::ops::IndexMut;
use std::sync::Arc;
use std::{fmt, usize};
use crate::physical_plan::joins::utils::{JoinFilter, JoinSide};
use crate::physical_plan::ExecutionPlan;
use arrow::compute::concat_batches;
use arrow::datatypes::{ArrowNativeType, SchemaRef};
use arrow_array::builder::BooleanBufferBuilder;
use arrow_array::{ArrowPrimitiveType, NativeAdapter, PrimitiveArray, RecordBatch};
use datafusion_common::tree_node::{Transformed, TreeNode};
use datafusion_common::{DataFusionError, Result, ScalarValue};
use datafusion_physical_expr::expressions::Column;
use datafusion_physical_expr::intervals::{ExprIntervalGraph, Interval, IntervalBound};
use datafusion_physical_expr::utils::collect_columns;
use datafusion_physical_expr::{PhysicalExpr, PhysicalSortExpr};
use hashbrown::raw::RawTable;
use hashbrown::HashSet;
use parking_lot::Mutex;
pub struct JoinHashMap {
pub map: RawTable<(u64, u64)>,
pub next: Vec<u64>,
}
impl JoinHashMap {
pub(crate) fn with_capacity(capacity: usize) -> Self {
JoinHashMap {
map: RawTable::with_capacity(capacity),
next: vec![0; capacity],
}
}
}
pub trait JoinHashMapType {
type NextType: IndexMut<usize, Output = u64>;
fn extend_zero(&mut self, len: usize);
fn get_mut(&mut self) -> (&mut RawTable<(u64, u64)>, &mut Self::NextType);
fn get_map(&self) -> &RawTable<(u64, u64)>;
fn get_list(&self) -> &Self::NextType;
}
impl JoinHashMapType for JoinHashMap {
type NextType = Vec<u64>;
fn extend_zero(&mut self, _: usize) {}
fn get_mut(&mut self) -> (&mut RawTable<(u64, u64)>, &mut Self::NextType) {
(&mut self.map, &mut self.next)
}
fn get_map(&self) -> &RawTable<(u64, u64)> {
&self.map
}
fn get_list(&self) -> &Self::NextType {
&self.next
}
}
impl JoinHashMapType for PruningJoinHashMap {
type NextType = VecDeque<u64>;
fn extend_zero(&mut self, len: usize) {
self.next.resize(self.next.len() + len, 0)
}
fn get_mut(&mut self) -> (&mut RawTable<(u64, u64)>, &mut Self::NextType) {
(&mut self.map, &mut self.next)
}
fn get_map(&self) -> &RawTable<(u64, u64)> {
&self.map
}
fn get_list(&self) -> &Self::NextType {
&self.next
}
}
impl fmt::Debug for JoinHashMap {
fn fmt(&self, _f: &mut fmt::Formatter) -> fmt::Result {
Ok(())
}
}
pub struct PruningJoinHashMap {
pub map: RawTable<(u64, u64)>,
pub next: VecDeque<u64>,
}
impl PruningJoinHashMap {
pub(crate) fn with_capacity(capacity: usize) -> Self {
PruningJoinHashMap {
map: RawTable::with_capacity(capacity),
next: VecDeque::with_capacity(capacity),
}
}
pub(crate) fn shrink_if_necessary(&mut self, scale_factor: usize) {
let capacity = self.map.capacity();
if capacity > scale_factor * self.map.len() {
let new_capacity = (capacity * (scale_factor - 1)) / scale_factor;
self.map.shrink_to(new_capacity, |(hash, _)| *hash)
}
}
pub(crate) fn size(&self) -> usize {
self.map.allocation_info().1.size()
+ self.next.capacity() * std::mem::size_of::<u64>()
}
pub(crate) fn prune_hash_values(
&mut self,
prune_length: usize,
deleting_offset: u64,
shrink_factor: usize,
) -> Result<()> {
self.next.drain(0..prune_length);
let removable_keys = unsafe {
self.map
.iter()
.map(|bucket| bucket.as_ref())
.filter_map(|(hash, tail_index)| {
(*tail_index < prune_length as u64 + deleting_offset).then_some(*hash)
})
.collect::<Vec<_>>()
};
removable_keys.into_iter().for_each(|hash_value| {
self.map
.remove_entry(hash_value, |(hash, _)| hash_value == *hash);
});
self.shrink_if_necessary(shrink_factor);
Ok(())
}
}
fn check_filter_expr_contains_sort_information(
expr: &Arc<dyn PhysicalExpr>,
reference: &Arc<dyn PhysicalExpr>,
) -> bool {
expr.eq(reference)
|| expr
.children()
.iter()
.any(|e| check_filter_expr_contains_sort_information(e, reference))
}
pub fn map_origin_col_to_filter_col(
filter: &JoinFilter,
schema: &SchemaRef,
side: &JoinSide,
) -> Result<HashMap<Column, Column>> {
let filter_schema = filter.schema();
let mut col_to_col_map: HashMap<Column, Column> = HashMap::new();
for (filter_schema_index, index) in filter.column_indices().iter().enumerate() {
if index.side.eq(side) {
let main_field = schema.field(index.index);
let main_col = Column::new_with_schema(main_field.name(), schema.as_ref())?;
let filter_field = filter_schema.field(filter_schema_index);
let filter_col = Column::new(filter_field.name(), filter_schema_index);
col_to_col_map.insert(main_col, filter_col);
}
}
Ok(col_to_col_map)
}
pub fn convert_sort_expr_with_filter_schema(
side: &JoinSide,
filter: &JoinFilter,
schema: &SchemaRef,
sort_expr: &PhysicalSortExpr,
) -> Result<Option<Arc<dyn PhysicalExpr>>> {
let column_map = map_origin_col_to_filter_col(filter, schema, side)?;
let expr = sort_expr.expr.clone();
let expr_columns = collect_columns(&expr);
let all_columns_are_included =
expr_columns.iter().all(|col| column_map.contains_key(col));
if all_columns_are_included {
let converted_filter_expr = expr.transform_up(&|p| {
convert_filter_columns(p.as_ref(), &column_map).map(|transformed| {
match transformed {
Some(transformed) => Transformed::Yes(transformed),
None => Transformed::No(p),
}
})
})?;
if check_filter_expr_contains_sort_information(
filter.expression(),
&converted_filter_expr,
) {
return Ok(Some(converted_filter_expr));
}
}
Ok(None)
}
pub fn build_filter_input_order(
side: JoinSide,
filter: &JoinFilter,
schema: &SchemaRef,
order: &PhysicalSortExpr,
) -> Result<Option<SortedFilterExpr>> {
let opt_expr = convert_sort_expr_with_filter_schema(&side, filter, schema, order)?;
Ok(opt_expr.map(|filter_expr| SortedFilterExpr::new(order.clone(), filter_expr)))
}
fn convert_filter_columns(
input: &dyn PhysicalExpr,
column_map: &HashMap<Column, Column>,
) -> Result<Option<Arc<dyn PhysicalExpr>>> {
Ok(if let Some(col) = input.as_any().downcast_ref::<Column>() {
column_map.get(col).map(|c| Arc::new(c.clone()) as _)
} else {
None
})
}
#[derive(Default)]
pub struct IntervalCalculatorInnerState {
graph: Option<ExprIntervalGraph>,
sorted_exprs: Vec<Option<SortedFilterExpr>>,
calculated: bool,
}
impl Debug for IntervalCalculatorInnerState {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "Exprs({:?})", self.sorted_exprs)
}
}
pub fn build_filter_expression_graph(
interval_state: &Arc<Mutex<IntervalCalculatorInnerState>>,
left: &Arc<dyn ExecutionPlan>,
right: &Arc<dyn ExecutionPlan>,
filter: &JoinFilter,
) -> Result<(
Option<SortedFilterExpr>,
Option<SortedFilterExpr>,
Option<ExprIntervalGraph>,
)> {
let mut filter_state = interval_state.lock();
if !filter_state.calculated {
let join_sides = [JoinSide::Left, JoinSide::Right];
let children = [left, right];
for (join_side, child) in join_sides.iter().zip(children.iter()) {
let sorted_expr = child
.output_ordering()
.and_then(|orders| {
build_filter_input_order(
*join_side,
filter,
&child.schema(),
&orders[0],
)
.transpose()
})
.transpose()?;
filter_state.sorted_exprs.push(sorted_expr);
}
let sorted_exprs_size = filter_state.sorted_exprs.len();
let mut sorted_exprs = filter_state
.sorted_exprs
.iter_mut()
.flatten()
.collect::<Vec<_>>();
filter_state.graph = if sorted_exprs.len() == sorted_exprs_size {
let mut graph = ExprIntervalGraph::try_new(filter.expression().clone())?;
let filter_exprs = sorted_exprs
.iter()
.map(|sorted_expr| sorted_expr.filter_expr().clone())
.collect::<Vec<_>>();
let child_node_indices = graph.gather_node_indices(&filter_exprs);
for (sorted_expr, (_, index)) in
sorted_exprs.iter_mut().zip(child_node_indices.iter())
{
sorted_expr.set_node_index(*index);
}
Some(graph)
} else {
None
};
filter_state.calculated = true;
}
Ok((
filter_state.sorted_exprs[0].clone(),
filter_state.sorted_exprs[1].clone(),
filter_state.graph.as_ref().cloned(),
))
}
#[derive(Debug, Clone)]
pub struct SortedFilterExpr {
origin_sorted_expr: PhysicalSortExpr,
filter_expr: Arc<dyn PhysicalExpr>,
interval: Interval,
node_index: usize,
}
impl SortedFilterExpr {
pub fn new(
origin_sorted_expr: PhysicalSortExpr,
filter_expr: Arc<dyn PhysicalExpr>,
) -> Self {
Self {
origin_sorted_expr,
filter_expr,
interval: Interval::default(),
node_index: 0,
}
}
pub fn origin_sorted_expr(&self) -> &PhysicalSortExpr {
&self.origin_sorted_expr
}
pub fn filter_expr(&self) -> &Arc<dyn PhysicalExpr> {
&self.filter_expr
}
pub fn interval(&self) -> &Interval {
&self.interval
}
pub fn set_interval(&mut self, interval: Interval) {
self.interval = interval;
}
pub fn node_index(&self) -> usize {
self.node_index
}
pub fn set_node_index(&mut self, node_index: usize) {
self.node_index = node_index;
}
}
pub fn calculate_filter_expr_intervals(
build_input_buffer: &RecordBatch,
build_sorted_filter_expr: &mut SortedFilterExpr,
probe_batch: &RecordBatch,
probe_sorted_filter_expr: &mut SortedFilterExpr,
) -> Result<()> {
if build_input_buffer.num_rows() == 0 || probe_batch.num_rows() == 0 {
return Ok(());
}
update_filter_expr_interval(
&build_input_buffer.slice(0, 1),
build_sorted_filter_expr,
)?;
update_filter_expr_interval(
&probe_batch.slice(probe_batch.num_rows() - 1, 1),
probe_sorted_filter_expr,
)
}
pub fn update_filter_expr_interval(
batch: &RecordBatch,
sorted_expr: &mut SortedFilterExpr,
) -> Result<()> {
let array = sorted_expr
.origin_sorted_expr()
.expr
.evaluate(batch)?
.into_array(1);
let value = ScalarValue::try_from_array(&array, 0)?;
let unbounded = IntervalBound::make_unbounded(value.get_datatype())?;
let interval = if sorted_expr.origin_sorted_expr().options.descending {
Interval::new(unbounded, IntervalBound::new(value, false))
} else {
Interval::new(IntervalBound::new(value, false), unbounded)
};
sorted_expr.set_interval(interval);
Ok(())
}
pub fn get_pruning_anti_indices<T: ArrowPrimitiveType>(
prune_length: usize,
deleted_offset: usize,
visited_rows: &HashSet<usize>,
) -> PrimitiveArray<T>
where
NativeAdapter<T>: From<<T as ArrowPrimitiveType>::Native>,
{
let mut bitmap = BooleanBufferBuilder::new(prune_length);
bitmap.append_n(prune_length, false);
for v in 0..prune_length {
let row = v + deleted_offset;
bitmap.set_bit(v, visited_rows.contains(&row));
}
(0..prune_length)
.filter_map(|idx| (!bitmap.get_bit(idx)).then_some(T::Native::from_usize(idx)))
.collect()
}
pub fn get_pruning_semi_indices<T: ArrowPrimitiveType>(
prune_length: usize,
deleted_offset: usize,
visited_rows: &HashSet<usize>,
) -> PrimitiveArray<T>
where
NativeAdapter<T>: From<<T as ArrowPrimitiveType>::Native>,
{
let mut bitmap = BooleanBufferBuilder::new(prune_length);
bitmap.append_n(prune_length, false);
(0..prune_length).for_each(|v| {
let row = &(v + deleted_offset);
bitmap.set_bit(v, visited_rows.contains(row));
});
(0..prune_length)
.filter_map(|idx| (bitmap.get_bit(idx)).then_some(T::Native::from_usize(idx)))
.collect::<PrimitiveArray<T>>()
}
pub fn combine_two_batches(
output_schema: &SchemaRef,
left_batch: Option<RecordBatch>,
right_batch: Option<RecordBatch>,
) -> Result<Option<RecordBatch>> {
match (left_batch, right_batch) {
(Some(batch), None) | (None, Some(batch)) => {
Ok(Some(batch))
}
(Some(left_batch), Some(right_batch)) => {
concat_batches(output_schema, &[left_batch, right_batch])
.map_err(DataFusionError::ArrowError)
.map(Some)
}
(None, None) => {
Ok(None)
}
}
}
pub fn record_visited_indices<T: ArrowPrimitiveType>(
visited: &mut HashSet<usize>,
offset: usize,
indices: &PrimitiveArray<T>,
) {
for i in indices.values() {
visited.insert(i.as_usize() + offset);
}
}
#[cfg(test)]
pub mod tests {
use super::*;
use crate::physical_plan::{
expressions::Column,
expressions::PhysicalSortExpr,
joins::utils::{ColumnIndex, JoinFilter, JoinSide},
};
use arrow::compute::SortOptions;
use arrow::datatypes::{DataType, Field, Schema};
use datafusion_common::ScalarValue;
use datafusion_expr::Operator;
use datafusion_physical_expr::expressions::{binary, cast, col, lit};
use std::sync::Arc;
pub(crate) fn complicated_filter(
filter_schema: &Schema,
) -> Result<Arc<dyn PhysicalExpr>> {
let left_expr = binary(
cast(
binary(
col("0", filter_schema)?,
Operator::Plus,
col("1", filter_schema)?,
filter_schema,
)?,
filter_schema,
DataType::Int64,
)?,
Operator::Gt,
binary(
cast(col("2", filter_schema)?, filter_schema, DataType::Int64)?,
Operator::Plus,
lit(ScalarValue::Int64(Some(10))),
filter_schema,
)?,
filter_schema,
)?;
let right_expr = binary(
cast(
binary(
col("0", filter_schema)?,
Operator::Plus,
col("1", filter_schema)?,
filter_schema,
)?,
filter_schema,
DataType::Int64,
)?,
Operator::Lt,
binary(
cast(col("2", filter_schema)?, filter_schema, DataType::Int64)?,
Operator::Plus,
lit(ScalarValue::Int64(Some(100))),
filter_schema,
)?,
filter_schema,
)?;
binary(left_expr, Operator::And, right_expr, filter_schema)
}
#[test]
fn test_column_exchange() -> Result<()> {
let left_child_schema =
Schema::new(vec![Field::new("left_1", DataType::Int32, true)]);
let left_child_sort_expr = PhysicalSortExpr {
expr: col("left_1", &left_child_schema)?,
options: SortOptions::default(),
};
let right_child_schema = Schema::new(vec![
Field::new("right_1", DataType::Int32, true),
Field::new("right_2", DataType::Int32, true),
]);
let right_child_sort_expr = PhysicalSortExpr {
expr: binary(
col("right_1", &right_child_schema)?,
Operator::Plus,
col("right_2", &right_child_schema)?,
&right_child_schema,
)?,
options: SortOptions::default(),
};
let intermediate_schema = Schema::new(vec![
Field::new("filter_1", DataType::Int32, true),
Field::new("filter_2", DataType::Int32, true),
Field::new("filter_3", DataType::Int32, true),
]);
let filter_left = col("filter_1", &intermediate_schema)?;
let filter_right = binary(
col("filter_2", &intermediate_schema)?,
Operator::Plus,
col("filter_3", &intermediate_schema)?,
&intermediate_schema,
)?;
let filter_expr = binary(
filter_left.clone(),
Operator::Gt,
filter_right.clone(),
&intermediate_schema,
)?;
let column_indices = vec![
ColumnIndex {
index: 0,
side: JoinSide::Left,
},
ColumnIndex {
index: 0,
side: JoinSide::Right,
},
ColumnIndex {
index: 1,
side: JoinSide::Right,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
let left_sort_filter_expr = build_filter_input_order(
JoinSide::Left,
&filter,
&Arc::new(left_child_schema),
&left_child_sort_expr,
)?
.unwrap();
assert!(left_child_sort_expr.eq(left_sort_filter_expr.origin_sorted_expr()));
let right_sort_filter_expr = build_filter_input_order(
JoinSide::Right,
&filter,
&Arc::new(right_child_schema),
&right_child_sort_expr,
)?
.unwrap();
assert!(right_child_sort_expr.eq(right_sort_filter_expr.origin_sorted_expr()));
assert!(filter_left.eq(left_sort_filter_expr.filter_expr()));
assert!(filter_right.eq(right_sort_filter_expr.filter_expr()));
Ok(())
}
#[test]
fn test_column_collector() -> Result<()> {
let 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(&schema)?;
let columns = collect_columns(&filter_expr);
assert_eq!(columns.len(), 3);
Ok(())
}
#[test]
fn find_expr_inside_expr() -> Result<()> {
let 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(&schema)?;
let expr_1 = Arc::new(Column::new("gnz", 0)) as _;
assert!(!check_filter_expr_contains_sort_information(
&filter_expr,
&expr_1
));
let expr_2 = col("1", &schema)? as _;
assert!(check_filter_expr_contains_sort_information(
&filter_expr,
&expr_2
));
let expr_3 = cast(
binary(
col("0", &schema)?,
Operator::Plus,
col("1", &schema)?,
&schema,
)?,
&schema,
DataType::Int64,
)?;
assert!(check_filter_expr_contains_sort_information(
&filter_expr,
&expr_3
));
let expr_4 = Arc::new(Column::new("1", 42)) as _;
assert!(!check_filter_expr_contains_sort_information(
&filter_expr,
&expr_4,
));
Ok(())
}
#[test]
fn build_sorted_expr() -> Result<()> {
let left_schema = Schema::new(vec![
Field::new("la1", DataType::Int32, false),
Field::new("lb1", DataType::Int32, false),
Field::new("lc1", DataType::Int32, false),
Field::new("lt1", DataType::Int32, false),
Field::new("la2", DataType::Int32, false),
Field::new("la1_des", DataType::Int32, false),
]);
let right_schema = Schema::new(vec![
Field::new("ra1", DataType::Int32, false),
Field::new("rb1", DataType::Int32, false),
Field::new("rc1", DataType::Int32, false),
Field::new("rt1", DataType::Int32, false),
Field::new("ra2", DataType::Int32, false),
Field::new("ra1_des", DataType::Int32, false),
]);
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);
let left_schema = Arc::new(left_schema);
let right_schema = Arc::new(right_schema);
assert!(build_filter_input_order(
JoinSide::Left,
&filter,
&left_schema,
&PhysicalSortExpr {
expr: col("la1", left_schema.as_ref())?,
options: SortOptions::default(),
}
)?
.is_some());
assert!(build_filter_input_order(
JoinSide::Left,
&filter,
&left_schema,
&PhysicalSortExpr {
expr: col("lt1", left_schema.as_ref())?,
options: SortOptions::default(),
}
)?
.is_none());
assert!(build_filter_input_order(
JoinSide::Right,
&filter,
&right_schema,
&PhysicalSortExpr {
expr: col("ra1", right_schema.as_ref())?,
options: SortOptions::default(),
}
)?
.is_some());
assert!(build_filter_input_order(
JoinSide::Right,
&filter,
&right_schema,
&PhysicalSortExpr {
expr: col("rb1", right_schema.as_ref())?,
options: SortOptions::default(),
}
)?
.is_none());
Ok(())
}
#[test]
fn sorted_filter_expr_build() -> Result<()> {
let intermediate_schema = Schema::new(vec![
Field::new("0", DataType::Int32, true),
Field::new("1", DataType::Int32, true),
]);
let filter_expr = binary(
col("0", &intermediate_schema)?,
Operator::Minus,
col("1", &intermediate_schema)?,
&intermediate_schema,
)?;
let column_indices = vec![
ColumnIndex {
index: 0,
side: JoinSide::Left,
},
ColumnIndex {
index: 1,
side: JoinSide::Left,
},
];
let filter = JoinFilter::new(filter_expr, column_indices, intermediate_schema);
let schema = Schema::new(vec![
Field::new("a", DataType::Int32, false),
Field::new("b", DataType::Int32, false),
]);
let sorted = PhysicalSortExpr {
expr: binary(
col("a", &schema)?,
Operator::Plus,
col("b", &schema)?,
&schema,
)?,
options: SortOptions::default(),
};
let res = convert_sort_expr_with_filter_schema(
&JoinSide::Left,
&filter,
&Arc::new(schema),
&sorted,
)?;
assert!(res.is_none());
Ok(())
}
#[test]
fn test_shrink_if_necessary() {
let scale_factor = 4;
let mut join_hash_map = PruningJoinHashMap::with_capacity(100);
let data_size = 2000;
let deleted_part = 3 * data_size / 4;
for hash_value in 0..data_size {
join_hash_map.map.insert(
hash_value,
(hash_value, hash_value),
|(hash, _)| *hash,
);
}
assert_eq!(join_hash_map.map.len(), data_size as usize);
assert!(join_hash_map.map.capacity() >= data_size as usize);
for hash_value in 0..deleted_part {
join_hash_map
.map
.remove_entry(hash_value, |(hash, _)| hash_value == *hash);
}
assert_eq!(join_hash_map.map.len(), (data_size - deleted_part) as usize);
let old_capacity = join_hash_map.map.capacity();
join_hash_map.shrink_if_necessary(scale_factor);
let new_expected_capacity =
join_hash_map.map.capacity() * (scale_factor - 1) / scale_factor;
assert!(join_hash_map.map.capacity() >= new_expected_capacity);
assert!(join_hash_map.map.capacity() <= old_capacity);
}
}