@@ -20,12 +20,62 @@ use crate::expr::analysis::AnnotationFn;
2020use crate :: expr:: analysis:: Annotations ;
2121use crate :: expr:: analysis:: descendent_annotations;
2222use crate :: expr:: get_item;
23+ use crate :: expr:: is_root;
24+ use crate :: expr:: label_is_fallible;
25+ use crate :: expr:: label_null_sensitive;
2326use crate :: expr:: pack;
2427use crate :: expr:: root;
2528use crate :: expr:: traversal:: NodeExt ;
2629use crate :: expr:: traversal:: NodeRewriter ;
2730use crate :: expr:: traversal:: Transformed ;
2831use crate :: expr:: traversal:: TraversalOrder ;
32+ use crate :: scalar_fn:: is_negative_cost;
33+
34+ fn references_root ( expr : & Expression ) -> bool {
35+ is_root ( expr) || expr. children ( ) . iter ( ) . any ( references_root)
36+ }
37+
38+ /// Split expression into two parts:
39+ ///
40+ /// left is the optional outer part that we want to apply to array after
41+ /// canonicalizing.
42+ /// right is the optional inner part that we want to apply to array before
43+ /// canonicalizing.
44+ ///
45+ /// We want to push to array only if expression has a negative cost, is
46+ /// infallible and null-insensitive.
47+ ///
48+ /// TODO(myrrc): This is a specialized version of partition(), and we want to
49+ /// unify this with expression partitioning logic.
50+ pub fn split_expression_for_pushdown ( expr : Expression ) -> ( Option < Expression > , Option < Expression > ) {
51+ let labelled_expr = expr. clone ( ) ;
52+ let fallible = label_is_fallible ( & labelled_expr) ;
53+ let null_sensitive = label_null_sensitive ( & labelled_expr) ;
54+ let mut inner: Option < Expression > = None ;
55+
56+ let outer = expr
57+ . transform_down ( |node| {
58+ if is_negative_cost ( node. id ( ) )
59+ && references_root ( & node)
60+ && !fallible. get ( & node) . copied ( ) . unwrap_or ( true )
61+ && !null_sensitive. get ( & node) . copied ( ) . unwrap_or ( true )
62+ {
63+ inner = Some ( node) ;
64+ Ok ( Transformed {
65+ value : root ( ) ,
66+ changed : true ,
67+ order : TraversalOrder :: Skip ,
68+ } )
69+ } else {
70+ Ok ( Transformed :: no ( node) )
71+ }
72+ } )
73+ . vortex_expect ( "infallible" )
74+ . into_inner ( ) ;
75+
76+ let outer = ( !is_root ( & outer) ) . then_some ( outer) ;
77+ ( outer, inner)
78+ }
2979
3080/// Partition an expression into sub-expressions that are uniquely associated with an annotation.
3181/// A root expression is also returned that can be used to recombine the results of the partitions
@@ -205,10 +255,19 @@ where
205255mod tests {
206256 use rstest:: fixture;
207257 use rstest:: rstest;
208-
258+ use vortex_array:: expr:: Expression ;
259+ use vortex_array:: expr:: byte_length;
260+ use vortex_array:: expr:: cast;
261+ use vortex_array:: expr:: like;
262+ use vortex_array:: expr:: traversal:: NodeExt ;
263+ use vortex_array:: expr:: traversal:: Transformed ;
264+ use vortex_array:: expr:: traversal:: TraversalOrder ;
265+
266+ use super :: split_expression_for_pushdown;
209267 use super :: * ;
210268 use crate :: dtype:: DType ;
211269 use crate :: dtype:: Nullability :: NonNullable ;
270+ use crate :: dtype:: PType ;
212271 use crate :: dtype:: PType :: I32 ;
213272 use crate :: dtype:: StructFields ;
214273 use crate :: expr:: analysis:: make_free_field_annotator;
@@ -348,4 +407,61 @@ mod tests {
348407 let expected_b = pack ( [ ( "b_0" , pack ( [ ( "b" , col ( "b" ) ) ] , NonNullable ) ) ] , NonNullable ) ;
349408 assert_eq ! ( part_b, & expected_b, "{part_b} {expected_b}" ) ;
350409 }
410+
411+ fn join_split_expr ( initial : & Expression , outer : Option < Expression > , inner : Option < Expression > ) {
412+ let outer_expr = outer. unwrap_or_else ( root) ;
413+ let inner_expr = inner. unwrap_or_else ( root) ;
414+ let expected = outer_expr
415+ . transform_down ( |node| {
416+ if !is_root ( & node) {
417+ return Ok ( Transformed :: no ( node) ) ;
418+ }
419+ Ok ( Transformed {
420+ value : inner_expr. clone ( ) ,
421+ changed : true ,
422+ order : TraversalOrder :: Skip ,
423+ } )
424+ } )
425+ . vortex_expect ( "infallible" ) ;
426+ assert_eq ! ( & expected. into_inner( ) , initial) ;
427+ }
428+
429+ #[ test]
430+ fn split_expr_cast_root ( ) {
431+ let ( outer, inner) = split_expression_for_pushdown ( root ( ) ) ;
432+ assert_eq ! ( outer, None ) ;
433+ assert_eq ! ( inner, None ) ; // Applying root to array is useless work
434+ }
435+
436+ #[ test]
437+ fn split_expr_partial_pushdown ( ) {
438+ let dtype = DType :: Primitive ( PType :: U64 , NonNullable ) ;
439+ let expr = cast ( byte_length ( root ( ) ) , dtype. clone ( ) ) ;
440+ let ( outer, inner) = split_expression_for_pushdown ( expr. clone ( ) ) ;
441+ // [0] = cast([1], dtype)
442+ // [1] = byte_length(root)
443+ assert_eq ! ( outer, Some ( cast( root( ) , dtype) ) ) ;
444+ assert_eq ! ( inner, Some ( byte_length( root( ) ) ) ) ;
445+ join_split_expr ( & expr, outer, inner) ;
446+ }
447+
448+ #[ test]
449+ fn split_expr_full_pushdown ( ) {
450+ let expr = byte_length ( root ( ) ) ;
451+ let ( outer, inner) = split_expression_for_pushdown ( expr. clone ( ) ) ;
452+ assert_eq ! ( outer, None ) ;
453+ assert_eq ! ( inner, Some ( byte_length( root( ) ) ) ) ;
454+ join_split_expr ( & expr, outer, inner) ;
455+ }
456+
457+ #[ test]
458+ fn split_expr_no_pushdown ( ) {
459+ // We can push down lit(), but it we replace
460+ // lit() with root(), the semantics change.
461+ let expr = like ( root ( ) , lit ( 1u64 ) ) ;
462+ let ( outer, inner) = split_expression_for_pushdown ( expr. clone ( ) ) ;
463+ assert_eq ! ( outer, Some ( expr. clone( ) ) ) ;
464+ assert_eq ! ( inner, None ) ;
465+ join_split_expr ( & expr, outer, inner) ;
466+ }
351467}
0 commit comments