@@ -28,14 +28,16 @@ use datafusion_common::config::ConfigOptions;
2828use datafusion_common:: error:: Result ;
2929use datafusion_common:: tree_node:: { Transformed , TransformedResult , TreeNode } ;
3030use datafusion_common:: { internal_err, JoinSide , JoinType } ;
31+ use datafusion_expr:: Operator ;
3132use datafusion_expr_common:: sort_properties:: SortProperties ;
32- use datafusion_physical_expr:: expressions:: Column ;
33+ use datafusion_physical_expr:: expressions:: { BinaryExpr , Column } ;
3334use datafusion_physical_expr:: LexOrdering ;
35+ use datafusion_physical_expr:: PhysicalExprRef ;
3436use datafusion_physical_plan:: execution_plan:: EmissionType ;
35- use datafusion_physical_plan:: joins:: utils:: ColumnIndex ;
37+ use datafusion_physical_plan:: joins:: utils:: { ColumnIndex , JoinFilter , JoinOn } ;
3638use datafusion_physical_plan:: joins:: {
37- CrossJoinExec , HashJoinExec , NestedLoopJoinExec , PartitionMode ,
38- StreamJoinPartitionMode , SymmetricHashJoinExec ,
39+ AsOfJoinCondition , AsOfJoinExec , CrossJoinExec , HashJoinExec , NestedLoopJoinExec ,
40+ PartitionMode , StreamJoinPartitionMode , SymmetricHashJoinExec ,
3941} ;
4042use datafusion_physical_plan:: { ExecutionPlan , ExecutionPlanProperties } ;
4143use 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.
321544pub type PipelineFixerSubrule =
322545 dyn Fn ( Arc < dyn ExecutionPlan > , & ConfigOptions ) -> Result < Arc < dyn ExecutionPlan > > ;
0 commit comments