diff --git a/vortex-array/src/expr/expression.rs b/vortex-array/src/expr/expression.rs index d7f85825dbe..e53474ad040 100644 --- a/vortex-array/src/expr/expression.rs +++ b/vortex-array/src/expr/expression.rs @@ -16,7 +16,10 @@ use vortex_session::VortexSession; use crate::dtype::DType; use crate::expr::display::DisplayTreeExpr; +use crate::expr::traversal::TraversalOrder; +use crate::expr::traversal::pre_order_visit_down; use crate::scalar_fn::ScalarFnRef; +use crate::scalar_fn::ScalarFnVTable; use crate::scalar_fn::fns::root::Root; use crate::stats::rewrite::StatsRewriteCtx; @@ -204,6 +207,30 @@ impl Expression { pub fn display_tree(&self) -> impl Display { DisplayTreeExpr(self) } + + /// Returns true if this expression contains expression E inside. + /// + /// # Example + /// + /// ```rust + /// # use vortex_array::scalar_fn::fns::literal::Literal; + /// # use vortex_array::expr::{eq, lit, root}; + /// let expression = &eq(root(), lit(3u64)); + /// assert!(expression.contains::().unwrap()); + /// let expression = root(); + /// assert!(!expression.contains::().unwrap()); + /// ``` + pub fn contains(&self) -> VortexResult { + let mut contains = false; + pre_order_visit_down(self, |node| { + if node.is::() { + contains = true; + return Ok(TraversalOrder::Stop); + } + Ok(TraversalOrder::Continue) + })?; + Ok(contains) + } } /// The default display implementation for expressions uses the 'SQL'-style format. diff --git a/vortex-array/src/expr/mod.rs b/vortex-array/src/expr/mod.rs index df7e3f59b97..2ae61509ceb 100644 --- a/vortex-array/src/expr/mod.rs +++ b/vortex-array/src/expr/mod.rs @@ -156,6 +156,10 @@ mod tests { use std::collections::hash_map::RandomState; use std::hash::BuildHasher; + use vortex_array::expr::eq; + use vortex_array::expr::lit; + use vortex_array::expr::root; + use super::*; use crate::dtype::DType; use crate::dtype::FieldNames; @@ -164,20 +168,18 @@ mod tests { use crate::dtype::StructFields; use crate::expr::and; use crate::expr::col; - use crate::expr::eq; use crate::expr::get_item; use crate::expr::gt; use crate::expr::gt_eq; - use crate::expr::lit; use crate::expr::lt; use crate::expr::lt_eq; use crate::expr::not; use crate::expr::not_eq; use crate::expr::or; - use crate::expr::root; use crate::expr::select; use crate::expr::select_exclude; use crate::scalar::Scalar; + use crate::scalar_fn::fns::literal::Literal; #[test] fn basic_expr_split_test() { @@ -311,4 +313,12 @@ mod tests { "{dog: 32u32, cat: \"rufus\"}" ); } + + #[test] + fn expr_contains() { + let expression = &eq(root(), lit(3u64)); + assert!(expression.contains::().unwrap()); + let expression = root(); + assert!(!expression.contains::().unwrap()); + } } diff --git a/vortex-layout/src/scan/scan_builder.rs b/vortex-layout/src/scan/scan_builder.rs index 03f8c49649d..abce5047109 100644 --- a/vortex-layout/src/scan/scan_builder.rs +++ b/vortex-layout/src/scan/scan_builder.rs @@ -39,6 +39,7 @@ use vortex_utils::parallelism::get_available_parallelism; use crate::LayoutReader; use crate::LayoutReaderRef; +use crate::layouts::row_idx::RowIdx; use crate::layouts::row_idx::RowIdxLayoutReader; use crate::scan::repeated_scan::RepeatedScan; use crate::scan::split_by::SplitBy; @@ -273,14 +274,20 @@ impl ScanBuilder { // conjunction splitting if a filter is provided. let mut layout_reader = self.layout_reader; - // Enrich the layout reader to support RowIdx expressions. + // Enrich the layout reader to support RowIdx expressions if scan uses #row_idx. // Note that this is applied below the filter layout reader since it can perform // better over individual conjunctions. - layout_reader = Arc::new(RowIdxLayoutReader::new( - self.row_offset, - layout_reader, - self.session.clone(), - )); + let mut found_row_idx = self.projection.contains::()?; + if !found_row_idx && let Some(filter) = self.filter.as_ref() { + found_row_idx = filter.contains::()?; + } + if found_row_idx { + layout_reader = Arc::new(RowIdxLayoutReader::new( + self.row_offset, + layout_reader, + self.session.clone(), + )); + } // Normalize and simplify the expressions. let projection = self.projection.optimize_recursive(layout_reader.dtype())?;