Skip to content

Commit 5398171

Browse files
committed
asof first iteration
1 parent 6eb90fd commit 5398171

3 files changed

Lines changed: 1409 additions & 49 deletions

File tree

datafusion/physical-optimizer/src/join_selection.rs

Lines changed: 272 additions & 49 deletions
Original file line numberDiff line numberDiff line change
@@ -28,14 +28,16 @@ use datafusion_common::config::ConfigOptions;
2828
use datafusion_common::error::Result;
2929
use datafusion_common::tree_node::{Transformed, TransformedResult, TreeNode};
3030
use datafusion_common::{internal_err, JoinSide, JoinType};
31+
use datafusion_expr::Operator;
3132
use datafusion_expr_common::sort_properties::SortProperties;
32-
use datafusion_physical_expr::expressions::Column;
33+
use datafusion_physical_expr::expressions::{BinaryExpr, Column};
3334
use datafusion_physical_expr::LexOrdering;
35+
use datafusion_physical_expr::PhysicalExprRef;
3436
use datafusion_physical_plan::execution_plan::EmissionType;
35-
use datafusion_physical_plan::joins::utils::ColumnIndex;
37+
use datafusion_physical_plan::joins::utils::{ColumnIndex, JoinFilter, JoinOn};
3638
use datafusion_physical_plan::joins::{
37-
CrossJoinExec, HashJoinExec, NestedLoopJoinExec, PartitionMode,
38-
StreamJoinPartitionMode, SymmetricHashJoinExec,
39+
AsOfJoinCondition, AsOfJoinExec, CrossJoinExec, HashJoinExec, NestedLoopJoinExec,
40+
PartitionMode, StreamJoinPartitionMode, SymmetricHashJoinExec,
3941
};
4042
use datafusion_physical_plan::{ExecutionPlan, ExecutionPlanProperties};
4143
use std::sync::Arc;
@@ -256,59 +258,60 @@ fn statistical_join_selection_subrule(
256258
collect_threshold_byte_size: usize,
257259
collect_threshold_num_rows: usize,
258260
) -> Result<Transformed<Arc<dyn ExecutionPlan>>> {
259-
let transformed =
260-
if let Some(hash_join) = plan.as_any().downcast_ref::<HashJoinExec>() {
261-
match hash_join.partition_mode() {
262-
PartitionMode::Auto => try_collect_left(
263-
hash_join,
264-
false,
265-
collect_threshold_byte_size,
266-
collect_threshold_num_rows,
267-
)?
261+
let transformed = if let Some(asof_join) = try_asof_join(&plan)? {
262+
Some(asof_join)
263+
} else if let Some(hash_join) = plan.as_any().downcast_ref::<HashJoinExec>() {
264+
match hash_join.partition_mode() {
265+
PartitionMode::Auto => try_collect_left(
266+
hash_join,
267+
false,
268+
collect_threshold_byte_size,
269+
collect_threshold_num_rows,
270+
)?
271+
.map_or_else(
272+
|| partitioned_hash_join(hash_join).map(Some),
273+
|v| Ok(Some(v)),
274+
)?,
275+
PartitionMode::CollectLeft => try_collect_left(hash_join, true, 0, 0)?
268276
.map_or_else(
269277
|| partitioned_hash_join(hash_join).map(Some),
270278
|v| Ok(Some(v)),
271279
)?,
272-
PartitionMode::CollectLeft => try_collect_left(hash_join, true, 0, 0)?
273-
.map_or_else(
274-
|| partitioned_hash_join(hash_join).map(Some),
275-
|v| Ok(Some(v)),
276-
)?,
277-
PartitionMode::Partitioned => {
278-
let left = hash_join.left();
279-
let right = hash_join.right();
280-
if hash_join.join_type().supports_swap()
281-
&& should_swap_join_order(&**left, &**right)?
282-
{
283-
hash_join
284-
.swap_inputs(PartitionMode::Partitioned)
285-
.map(Some)?
286-
} else {
287-
None
288-
}
280+
PartitionMode::Partitioned => {
281+
let left = hash_join.left();
282+
let right = hash_join.right();
283+
if hash_join.join_type().supports_swap()
284+
&& should_swap_join_order(&**left, &**right)?
285+
{
286+
hash_join
287+
.swap_inputs(PartitionMode::Partitioned)
288+
.map(Some)?
289+
} else {
290+
None
289291
}
290292
}
291-
} else if let Some(cross_join) = plan.as_any().downcast_ref::<CrossJoinExec>() {
292-
let left = cross_join.left();
293-
let right = cross_join.right();
294-
if should_swap_join_order(&**left, &**right)? {
295-
cross_join.swap_inputs().map(Some)?
296-
} else {
297-
None
298-
}
299-
} else if let Some(nl_join) = plan.as_any().downcast_ref::<NestedLoopJoinExec>() {
300-
let left = nl_join.left();
301-
let right = nl_join.right();
302-
if nl_join.join_type().supports_swap()
303-
&& should_swap_join_order(&**left, &**right)?
304-
{
305-
nl_join.swap_inputs().map(Some)?
306-
} else {
307-
None
308-
}
293+
}
294+
} else if let Some(cross_join) = plan.as_any().downcast_ref::<CrossJoinExec>() {
295+
let left = cross_join.left();
296+
let right = cross_join.right();
297+
if should_swap_join_order(&**left, &**right)? {
298+
cross_join.swap_inputs().map(Some)?
309299
} else {
310300
None
311-
};
301+
}
302+
} else if let Some(nl_join) = plan.as_any().downcast_ref::<NestedLoopJoinExec>() {
303+
let left = nl_join.left();
304+
let right = nl_join.right();
305+
if nl_join.join_type().supports_swap()
306+
&& should_swap_join_order(&**left, &**right)?
307+
{
308+
nl_join.swap_inputs().map(Some)?
309+
} else {
310+
None
311+
}
312+
} else {
313+
None
314+
};
312315

313316
Ok(if let Some(transformed) = transformed {
314317
Transformed::yes(transformed)
@@ -317,6 +320,226 @@ fn statistical_join_selection_subrule(
317320
})
318321
}
319322

323+
fn try_asof_join(
324+
plan: &Arc<dyn ExecutionPlan>,
325+
) -> Result<Option<Arc<dyn ExecutionPlan>>> {
326+
if let Some(hash_join) = plan.as_any().downcast_ref::<HashJoinExec>() {
327+
if !matches!(hash_join.join_type(), JoinType::Inner | JoinType::Left) {
328+
return Ok(None);
329+
}
330+
331+
let Some(filter) = hash_join.filter() else {
332+
return Ok(None);
333+
};
334+
let Some((filter_on, asof_condition)) =
335+
extract_asof_join_predicates(filter, hash_join.left(), hash_join.right())?
336+
else {
337+
return Ok(None);
338+
};
339+
340+
let mut on = hash_join.on().to_vec();
341+
on.extend(filter_on);
342+
343+
return Ok(Some(Arc::new(AsOfJoinExec::try_new(
344+
Arc::clone(hash_join.left()),
345+
Arc::clone(hash_join.right()),
346+
on,
347+
asof_condition,
348+
*hash_join.join_type(),
349+
hash_join.projection.clone(),
350+
hash_join.null_equality(),
351+
)?)));
352+
}
353+
354+
if let Some(nl_join) = plan.as_any().downcast_ref::<NestedLoopJoinExec>() {
355+
if !matches!(nl_join.join_type(), JoinType::Inner | JoinType::Left) {
356+
return Ok(None);
357+
}
358+
359+
let Some(filter) = nl_join.filter() else {
360+
return Ok(None);
361+
};
362+
let Some((on, asof_condition)) =
363+
extract_asof_join_predicates(filter, nl_join.left(), nl_join.right())?
364+
else {
365+
return Ok(None);
366+
};
367+
368+
return Ok(Some(Arc::new(AsOfJoinExec::try_new(
369+
Arc::clone(nl_join.left()),
370+
Arc::clone(nl_join.right()),
371+
on,
372+
asof_condition,
373+
*nl_join.join_type(),
374+
nl_join.projection().cloned(),
375+
datafusion_common::NullEquality::NullEqualsNothing,
376+
)?)));
377+
}
378+
379+
Ok(None)
380+
}
381+
382+
fn extract_asof_join_predicates(
383+
filter: &JoinFilter,
384+
left: &Arc<dyn ExecutionPlan>,
385+
right: &Arc<dyn ExecutionPlan>,
386+
) -> Result<Option<(JoinOn, AsOfJoinCondition)>> {
387+
let mut visitor = AsOfPredicateVisitor {
388+
filter,
389+
left,
390+
right,
391+
on: vec![],
392+
asof_condition: None,
393+
inequality_count: 0,
394+
unsupported: false,
395+
};
396+
visitor.visit(filter.expression())?;
397+
398+
if visitor.unsupported || visitor.inequality_count != 1 {
399+
return Ok(None);
400+
}
401+
402+
Ok(visitor
403+
.asof_condition
404+
.map(|asof_condition| (visitor.on, asof_condition)))
405+
}
406+
407+
struct AsOfPredicateVisitor<'a> {
408+
filter: &'a JoinFilter,
409+
left: &'a Arc<dyn ExecutionPlan>,
410+
right: &'a Arc<dyn ExecutionPlan>,
411+
on: JoinOn,
412+
asof_condition: Option<AsOfJoinCondition>,
413+
inequality_count: usize,
414+
unsupported: bool,
415+
}
416+
417+
impl AsOfPredicateVisitor<'_> {
418+
fn visit(&mut self, expr: &PhysicalExprRef) -> Result<()> {
419+
let Some(binary) = expr.as_any().downcast_ref::<BinaryExpr>() else {
420+
self.unsupported = true;
421+
return Ok(());
422+
};
423+
424+
if binary.op() == &Operator::And {
425+
self.visit(binary.left())?;
426+
self.visit(binary.right())?;
427+
return Ok(());
428+
}
429+
430+
match binary.op() {
431+
Operator::Eq => {
432+
if let Some((left, right)) =
433+
self.extract_side_pair(binary.left(), binary.right(), *binary.op())?
434+
{
435+
self.on.push((left, right));
436+
} else {
437+
self.unsupported = true;
438+
}
439+
}
440+
Operator::Lt | Operator::LtEq | Operator::Gt | Operator::GtEq => {
441+
self.inequality_count += 1;
442+
if self.inequality_count > 1 {
443+
self.unsupported = true;
444+
return Ok(());
445+
}
446+
447+
if let Some((left, right, op)) =
448+
self.extract_inequality(binary.left(), binary.right(), *binary.op())?
449+
{
450+
self.asof_condition =
451+
Some(AsOfJoinCondition::try_new(left, op, right)?);
452+
} else {
453+
self.unsupported = true;
454+
}
455+
}
456+
_ => self.unsupported = true,
457+
}
458+
459+
Ok(())
460+
}
461+
462+
fn extract_side_pair(
463+
&self,
464+
left_expr: &PhysicalExprRef,
465+
right_expr: &PhysicalExprRef,
466+
op: Operator,
467+
) -> Result<Option<(PhysicalExprRef, PhysicalExprRef)>> {
468+
let Some((left_side, left_expr)) = self.filter_column(left_expr)? else {
469+
return Ok(None);
470+
};
471+
let Some((right_side, right_expr)) = self.filter_column(right_expr)? else {
472+
return Ok(None);
473+
};
474+
475+
match (left_side, right_side, op) {
476+
(JoinSide::Left, JoinSide::Right, Operator::Eq) => {
477+
Ok(Some((left_expr, right_expr)))
478+
}
479+
(JoinSide::Right, JoinSide::Left, Operator::Eq) => {
480+
Ok(Some((right_expr, left_expr)))
481+
}
482+
_ => Ok(None),
483+
}
484+
}
485+
486+
fn extract_inequality(
487+
&self,
488+
left_expr: &PhysicalExprRef,
489+
right_expr: &PhysicalExprRef,
490+
op: Operator,
491+
) -> Result<Option<(PhysicalExprRef, PhysicalExprRef, Operator)>> {
492+
let Some((left_side, left_expr)) = self.filter_column(left_expr)? else {
493+
return Ok(None);
494+
};
495+
let Some((right_side, right_expr)) = self.filter_column(right_expr)? else {
496+
return Ok(None);
497+
};
498+
499+
match (left_side, right_side) {
500+
(JoinSide::Left, JoinSide::Right) => Ok(Some((left_expr, right_expr, op))),
501+
(JoinSide::Right, JoinSide::Left) => {
502+
Ok(Some((right_expr, left_expr, flip_inequality(op)?)))
503+
}
504+
_ => Ok(None),
505+
}
506+
}
507+
508+
fn filter_column(
509+
&self,
510+
expr: &PhysicalExprRef,
511+
) -> Result<Option<(JoinSide, PhysicalExprRef)>> {
512+
let Some(column) = expr.as_any().downcast_ref::<Column>() else {
513+
return Ok(None);
514+
};
515+
let Some(column_index) = self.filter.column_indices().get(column.index()) else {
516+
return Ok(None);
517+
};
518+
519+
let (schema, side) = match column_index.side {
520+
JoinSide::Left => (self.left.schema(), JoinSide::Left),
521+
JoinSide::Right => (self.right.schema(), JoinSide::Right),
522+
JoinSide::None => return Ok(None),
523+
};
524+
525+
let field = schema.field(column_index.index);
526+
Ok(Some((
527+
side,
528+
Arc::new(Column::new(field.name(), column_index.index)) as _,
529+
)))
530+
}
531+
}
532+
533+
fn flip_inequality(op: Operator) -> Result<Operator> {
534+
match op {
535+
Operator::Lt => Ok(Operator::Gt),
536+
Operator::LtEq => Ok(Operator::GtEq),
537+
Operator::Gt => Ok(Operator::Lt),
538+
Operator::GtEq => Ok(Operator::LtEq),
539+
_ => internal_err!("Can not flip non-inequality operator {op}"),
540+
}
541+
}
542+
320543
/// Pipeline-fixing join selection subrule.
321544
pub type PipelineFixerSubrule =
322545
dyn Fn(Arc<dyn ExecutionPlan>, &ConfigOptions) -> Result<Arc<dyn ExecutionPlan>>;

0 commit comments

Comments
 (0)