diff --git a/vortex-array/src/scalar_fn/fns/select.rs b/vortex-array/src/scalar_fn/fns/select.rs index b0ac654a2f1..9d74614be70 100644 --- a/vortex-array/src/scalar_fn/fns/select.rs +++ b/vortex-array/src/scalar_fn/fns/select.rs @@ -19,6 +19,8 @@ use vortex_session::registry::CachedId; use crate::ArrayRef; use crate::ExecutionCtx; use crate::IntoArray; +use crate::arrays::Constant; +use crate::arrays::ConstantArray; use crate::arrays::StructArray; use crate::arrays::struct_::StructArrayExt; use crate::dtype::DType; @@ -149,7 +151,19 @@ impl ScalarFnVTable for Select { args: &dyn ExecutionArgs, ctx: &mut ExecutionCtx, ) -> VortexResult { - let child = args.get(0)?.execute::(ctx)?; + let child = args.get(0)?; + if let Some(constant) = child.as_opt::() { + let child_struct_dtype = child + .dtype() + .as_struct_fields_opt() + .ok_or_else(|| vortex_err!("Select child not a struct dtype"))?; + let included = selection.normalize_to_included_fields(child_struct_dtype.names())?; + let scalar = constant.scalar().as_struct().project(included.as_ref())?; + + return Ok(ConstantArray::new(scalar, child.len()).into_array()); + } + + let child = child.execute::(ctx)?; let result = match selection { FieldSelection::Include(f) => child.project(f.as_ref()), @@ -308,10 +322,15 @@ impl Display for FieldSelection { #[cfg(test)] mod tests { use vortex_buffer::buffer; + use vortex_error::VortexResult; + use crate::ArrayRef; use crate::IntoArray; use crate::VortexSessionExecute; use crate::array_session; + use crate::arrays::Constant; + use crate::arrays::ConstantArray; + use crate::arrays::ScalarFnArray; use crate::arrays::struct_::StructArrayExt; use crate::dtype::DType; use crate::dtype::FieldName; @@ -324,6 +343,9 @@ mod tests { use crate::expr::select; use crate::expr::select_exclude; use crate::expr::test_harness; + use crate::scalar::Scalar; + use crate::scalar_fn::ScalarFnVTableExt; + use crate::scalar_fn::fns::select::FieldSelection; use crate::scalar_fn::fns::select::Select; use crate::scalar_fn::fns::select::StructArray; @@ -428,6 +450,46 @@ mod tests { ); } + #[test] + fn execute_constant_struct_stays_constant() -> VortexResult<()> { + let field_count = 128usize; + let names = (0..field_count) + .map(|idx| FieldName::from(format!("f{idx}"))) + .collect::>(); + let dtypes = vec![I32.into(); field_count]; + let fields = StructFields::new(FieldNames::from(names), dtypes); + let scalar = Scalar::struct_( + DType::Struct(fields, Nullability::NonNullable), + (0..128).map(Scalar::from), + ); + let array = ConstantArray::new(scalar, 1_000_000).into_array(); + let mut ctx = array_session().create_execution_ctx(); + + let item: ArrayRef = ScalarFnArray::try_new( + Select.bind(FieldSelection::Include(FieldNames::from(["f97", "f3"]))), + vec![array], + )? + .into_array() + .execute(&mut ctx)?; + + let constant = item.as_::(); + assert_eq!(constant.len(), 1_000_000); + assert_eq!( + constant.scalar(), + &Scalar::struct_( + DType::Struct( + StructFields::new( + FieldNames::from(["f97", "f3"]), + vec![I32.into(), I32.into()], + ), + Nullability::NonNullable, + ), + [Scalar::from(97), Scalar::from(3)], + ) + ); + Ok(()) + } + #[test] fn test_remove_select_rule() { let dtype = DType::Struct(