diff --git a/vortex-array/src/arrays/filter/rules.rs b/vortex-array/src/arrays/filter/rules.rs index 635ae198fca..53f35c154cd 100644 --- a/vortex-array/src/arrays/filter/rules.rs +++ b/vortex-array/src/arrays/filter/rules.rs @@ -15,13 +15,20 @@ use crate::arrays::filter::FilterArraySlotsExt; use crate::arrays::filter::FilterReduce; use crate::arrays::filter::FilterReduceAdaptor; use crate::arrays::filter::execute::buffer::prepare_mask_for_reuse; +use crate::arrays::scalar_fn::ExactScalarFn; +use crate::arrays::scalar_fn::ScalarFnArrayView; use crate::arrays::struct_::StructDataParts; +use crate::builtins::ArrayBuiltins; +use crate::optimizer::rules::ArrayParentReduceRule; use crate::optimizer::rules::ArrayReduceRule; use crate::optimizer::rules::ParentRuleSet; use crate::optimizer::rules::ReduceRuleSet; +use crate::scalar_fn::fns::get_item::GetItem; -pub(super) const PARENT_RULES: ParentRuleSet = - ParentRuleSet::new(&[ParentRuleSet::lift(&FilterReduceAdaptor(Filter))]); +pub(super) const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ + ParentRuleSet::lift(&FilterReduceAdaptor(Filter)), + ParentRuleSet::lift(&FilterGetItemRule), +]); pub(super) const RULES: ReduceRuleSet = ReduceRuleSet::new(&[&TrivialFilterRule, &FilterStructRule]); @@ -35,6 +42,23 @@ impl FilterReduce for Filter { } } +#[derive(Debug)] +struct FilterGetItemRule; + +impl ArrayParentReduceRule for FilterGetItemRule { + type Parent = ExactScalarFn; + + fn reduce_parent( + &self, + array: ArrayView<'_, Filter>, + parent: ScalarFnArrayView<'_, GetItem>, + _child_idx: usize, + ) -> VortexResult> { + let field = array.child().get_item(parent.options.clone())?; + Ok(Some(field.filter(array.filter_mask().clone())?)) + } +} + #[derive(Debug)] struct TrivialFilterRule; diff --git a/vortex-array/src/arrays/masked/compute/rules.rs b/vortex-array/src/arrays/masked/compute/rules.rs index 3accb455c3f..ac2123aedc8 100644 --- a/vortex-array/src/arrays/masked/compute/rules.rs +++ b/vortex-array/src/arrays/masked/compute/rules.rs @@ -1,16 +1,47 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors +use vortex_error::VortexResult; + +use crate::ArrayRef; +use crate::array::ArrayView; use crate::arrays::Masked; use crate::arrays::dict::TakeReduceAdaptor; use crate::arrays::filter::FilterReduceAdaptor; +use crate::arrays::masked::MaskedArrayExt; +use crate::arrays::masked::MaskedArraySlotsExt; +use crate::arrays::scalar_fn::ExactScalarFn; +use crate::arrays::scalar_fn::ScalarFnArrayView; use crate::arrays::slice::SliceReduceAdaptor; +use crate::builtins::ArrayBuiltins; +use crate::optimizer::rules::ArrayParentReduceRule; use crate::optimizer::rules::ParentRuleSet; +use crate::scalar_fn::fns::get_item::GetItem; use crate::scalar_fn::fns::mask::MaskReduceAdaptor; pub(crate) const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ ParentRuleSet::lift(&FilterReduceAdaptor(Masked)), + ParentRuleSet::lift(&MaskedGetItemRule), ParentRuleSet::lift(&MaskReduceAdaptor(Masked)), ParentRuleSet::lift(&SliceReduceAdaptor(Masked)), ParentRuleSet::lift(&TakeReduceAdaptor(Masked)), ]); + +#[derive(Debug)] +struct MaskedGetItemRule; + +impl ArrayParentReduceRule for MaskedGetItemRule { + type Parent = ExactScalarFn; + + fn reduce_parent( + &self, + array: ArrayView<'_, Masked>, + parent: ScalarFnArrayView<'_, GetItem>, + _child_idx: usize, + ) -> VortexResult> { + let field = array.child().get_item(parent.options.clone())?; + Ok(Some( + field.mask(array.masked_validity().to_array(array.len()))?, + )) + } +} diff --git a/vortex-array/src/arrays/slice/rules.rs b/vortex-array/src/arrays/slice/rules.rs index 3bc0a70f204..62ed8892541 100644 --- a/vortex-array/src/arrays/slice/rules.rs +++ b/vortex-array/src/arrays/slice/rules.rs @@ -1,9 +1,38 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors +use vortex_error::VortexResult; + +use crate::ArrayRef; +use crate::array::ArrayView; use crate::arrays::Slice; +use crate::arrays::scalar_fn::ExactScalarFn; +use crate::arrays::scalar_fn::ScalarFnArrayView; +use crate::arrays::slice::SliceArraySlotsExt; use crate::arrays::slice::SliceReduceAdaptor; +use crate::builtins::ArrayBuiltins; +use crate::optimizer::rules::ArrayParentReduceRule; use crate::optimizer::rules::ParentRuleSet; +use crate::scalar_fn::fns::get_item::GetItem; + +pub(super) const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ + ParentRuleSet::lift(&SliceReduceAdaptor(Slice)), + ParentRuleSet::lift(&SliceGetItemRule), +]); + +#[derive(Debug)] +struct SliceGetItemRule; + +impl ArrayParentReduceRule for SliceGetItemRule { + type Parent = ExactScalarFn; -pub(super) const PARENT_RULES: ParentRuleSet = - ParentRuleSet::new(&[ParentRuleSet::lift(&SliceReduceAdaptor(Slice))]); + fn reduce_parent( + &self, + array: ArrayView<'_, Slice>, + parent: ScalarFnArrayView<'_, GetItem>, + _child_idx: usize, + ) -> VortexResult> { + let field = array.child().get_item(parent.options.clone())?; + Ok(Some(field.slice(array.slice_range().clone())?)) + } +} diff --git a/vortex-array/src/scalar_fn/fns/get_item.rs b/vortex-array/src/scalar_fn/fns/get_item.rs index cec7b98e442..3e78c01f76b 100644 --- a/vortex-array/src/scalar_fn/fns/get_item.rs +++ b/vortex-array/src/scalar_fn/fns/get_item.rs @@ -243,6 +243,11 @@ mod tests { use crate::VortexSessionExecute; use crate::arrays::Constant; use crate::arrays::ConstantArray; + use crate::arrays::Filter; + use crate::arrays::Primitive; + use crate::arrays::ScalarFn; + use crate::arrays::Slice; + use crate::builtins::ArrayBuiltins; use crate::dtype::DType; use crate::dtype::FieldName; use crate::dtype::FieldNames; @@ -354,6 +359,50 @@ mod tests { Ok(()) } + #[test] + fn get_item_pushes_through_filter() -> VortexResult<()> { + use vortex_mask::Mask; + + let filtered = test_array() + .into_array() + .filter(Mask::from_iter([true, false, true]))?; + + let item = filtered.get_item("a")?; + + assert!(item.is::()); + assert!(!item.is::()); + Ok(()) + } + + #[test] + fn get_item_pushes_through_slice() -> VortexResult<()> { + let sliced = test_array().into_array().slice(1..3)?; + + let item = sliced.get_item("a")?; + + assert!(item.is::()); + assert!(!item.is::()); + assert!(!item.is::()); + Ok(()) + } + + #[test] + fn get_item_pushes_through_masked() -> VortexResult<()> { + let masked = test_array() + .into_array() + .mask(Validity::from_iter([true, false, true]).to_array(3))?; + + let item = masked.get_item("a")?; + + assert!(item.is::()); + assert_eq!( + item.dtype(), + &DType::Primitive(PType::I32, Nullability::Nullable) + ); + assert!(!item.is::()); + Ok(()) + } + #[test] fn test_pack_get_item_rule() { // Create: pack(a: lit(1), b: lit(2)).get_item("b")