Skip to main content

vortex_array/arrays/scalar_fn/vtable/
operations.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use vortex_error::VortexResult;
5
6use crate::ExecutionCtx;
7use crate::IntoArray;
8use crate::array::ArrayView;
9use crate::array::OperationsVTable;
10use crate::arrays::ConstantArray;
11use crate::arrays::scalar_fn::ScalarFnArrayExt;
12use crate::arrays::scalar_fn::vtable::ScalarFn;
13use crate::columnar::Columnar;
14use crate::scalar::Scalar;
15use crate::scalar_fn::VecExecutionArgs;
16
17impl OperationsVTable<ScalarFn> for ScalarFn {
18    fn scalar_at(
19        array: ArrayView<'_, ScalarFn>,
20        index: usize,
21        ctx: &mut ExecutionCtx,
22    ) -> VortexResult<Scalar> {
23        let inputs: Vec<_> = array
24            .children()
25            .iter()
26            .map(|child| Ok(ConstantArray::new(child.execute_scalar(index, ctx)?, 1).into_array()))
27            .collect::<VortexResult<_>>()?;
28
29        let args = VecExecutionArgs::new(inputs, 1);
30        let result = array.scalar_fn().execute(&args, ctx)?;
31
32        let scalar = match result.execute::<Columnar>(ctx)? {
33            Columnar::Canonical(arr) => {
34                tracing::info!(
35                    "Scalar function {} returned non-constant array from execution over all scalar inputs",
36                    array.scalar_fn(),
37                );
38                arr.into_array().execute_scalar(0, ctx)?
39            }
40            Columnar::Constant(constant) => constant.scalar().clone(),
41        };
42
43        debug_assert_eq!(
44            scalar.dtype(),
45            array.dtype(),
46            "Scalar function {} returned dtype {:?} but expected {:?}",
47            array.scalar_fn(),
48            scalar.dtype(),
49            array.dtype()
50        );
51
52        Ok(scalar)
53    }
54}
55
56#[cfg(test)]
57mod tests {
58    use vortex_buffer::buffer;
59    use vortex_error::VortexResult;
60
61    use crate::ArraySlots;
62    use crate::Canonical;
63    use crate::IntoArray;
64    use crate::VortexSessionExecute;
65    use crate::array::Array;
66    use crate::array::ArrayParts;
67    use crate::array_session;
68    use crate::arrays::BoolArray;
69    use crate::arrays::PrimitiveArray;
70    use crate::arrays::ScalarFnArray;
71    use crate::arrays::scalar_fn::ScalarFnArrayExt;
72    use crate::arrays::scalar_fn::array::ScalarFnData;
73    use crate::arrays::scalar_fn::vtable::ScalarFn;
74    use crate::assert_arrays_eq;
75    use crate::scalar::Scalar;
76    use crate::scalar_fn::TypedScalarFnInstance;
77    use crate::scalar_fn::fns::binary::Binary;
78    use crate::scalar_fn::fns::literal::Literal;
79    use crate::scalar_fn::fns::operators::Operator;
80    use crate::validity::Validity;
81
82    #[test]
83    fn test_scalar_fn_add() -> VortexResult<()> {
84        let mut ctx = array_session().create_execution_ctx();
85        let lhs = buffer![1i32, 2, 3].into_array();
86        let rhs = buffer![10i32, 20, 30].into_array();
87
88        let scalar_fn = TypedScalarFnInstance::new(Binary, Operator::Add).erased();
89        let scalar_fn_array = ScalarFnArray::try_new(scalar_fn, vec![lhs, rhs])?;
90
91        assert_eq!(scalar_fn_array.len(), 3);
92
93        let result = scalar_fn_array
94            .into_array()
95            .execute::<Canonical>(&mut array_session().create_execution_ctx())?
96            .into_array();
97        let expected = buffer![11i32, 22, 33].into_array();
98        assert_arrays_eq!(result, expected, &mut ctx);
99
100        Ok(())
101    }
102
103    #[test]
104    fn test_scalar_fn_inferred_len_rejects_mismatched_children() {
105        let lhs = buffer![1i32, 2, 3].into_array();
106        let rhs = buffer![10i32, 20].into_array();
107
108        let scalar_fn = TypedScalarFnInstance::new(Binary, Operator::Add).erased();
109        let err = ScalarFnArray::try_new(scalar_fn, vec![lhs, rhs])
110            .expect_err("ScalarFnArray::try_new must reject mismatched child lengths");
111
112        assert!(
113            err.to_string()
114                .contains("ScalarFnArray must have children equal to the array length")
115        );
116    }
117
118    #[test]
119    fn rejects_wrong_arity() {
120        let child = buffer![1i32, 2, 3].into_array();
121        let scalar_fn = TypedScalarFnInstance::new(Binary, Operator::Add).erased();
122
123        let Err(err) = ScalarFnArray::try_new(scalar_fn.clone(), vec![child.clone()]) else {
124            panic!("Binary must reject one child");
125        };
126        assert!(
127            err.to_string()
128                .contains("ScalarFnArray requires 2 children, got 1")
129        );
130
131        let Err(err) = ScalarFnArray::try_new(scalar_fn, vec![child.clone(), child.clone(), child])
132        else {
133            panic!("Binary must reject three children");
134        };
135        assert!(
136            err.to_string()
137                .contains("ScalarFnArray requires 2 children, got 3")
138        );
139    }
140
141    #[test]
142    fn rejects_missing_child_slot() {
143        let lhs = buffer![1i32, 2, 3].into_array();
144        let dtype = lhs.dtype().clone();
145        let scalar_fn = TypedScalarFnInstance::new(Binary, Operator::Add).erased();
146        let vtable = ScalarFn { id: scalar_fn.id() };
147        let data = ScalarFnData { scalar_fn };
148        let slots = [Some(lhs), None].into_iter().collect::<ArraySlots>();
149
150        let Err(err) = Array::<ScalarFn>::try_from_parts(
151            ArrayParts::new(vtable, dtype, 3, data).with_slots(slots),
152        ) else {
153            panic!("ScalarFnArray must reject missing child slots");
154        };
155
156        assert!(
157            err.to_string()
158                .contains("ScalarFnArray requires every child slot to be present, got 1 missing")
159        );
160    }
161
162    #[test]
163    fn test_scalar_fn_without_children_requires_explicit_len() -> VortexResult<()> {
164        let scalar_fn = TypedScalarFnInstance::new(Literal, Scalar::from(1i32)).erased();
165
166        let Err(err) = ScalarFnArray::try_new(scalar_fn.clone(), vec![]) else {
167            panic!("ScalarFnArray::try_new should reject zero children");
168        };
169        assert!(
170            err.to_string()
171                .contains("ScalarFnArray length cannot be inferred without children")
172        );
173
174        let scalar_fn_array = ScalarFnArray::try_new_with_len(scalar_fn, vec![], 3)?;
175        assert_eq!(scalar_fn_array.len(), 3);
176        assert_eq!(scalar_fn_array.child_count(), 0);
177
178        Ok(())
179    }
180
181    #[test]
182    fn test_scalar_fn_mul() -> VortexResult<()> {
183        let mut ctx = array_session().create_execution_ctx();
184        let lhs = buffer![2i32, 3, 4].into_array();
185        let rhs = buffer![5i32, 6, 7].into_array();
186
187        let scalar_fn = TypedScalarFnInstance::new(Binary, Operator::Mul).erased();
188        let scalar_fn_array = ScalarFnArray::try_new(scalar_fn, vec![lhs, rhs])?;
189
190        let result = scalar_fn_array
191            .into_array()
192            .execute::<Canonical>(&mut array_session().create_execution_ctx())?
193            .into_array();
194        let expected = buffer![10i32, 18, 28].into_array();
195        assert_arrays_eq!(result, expected, &mut ctx);
196
197        Ok(())
198    }
199
200    #[test]
201    fn test_scalar_fn_with_nullable() -> VortexResult<()> {
202        let mut ctx = array_session().create_execution_ctx();
203        let lhs = PrimitiveArray::new(buffer![1i32, 2, 3], Validity::AllValid).into_array();
204        let rhs = PrimitiveArray::new(
205            buffer![10i32, 20, 30],
206            Validity::from_iter([true, false, true]),
207        )
208        .into_array();
209
210        let scalar_fn = TypedScalarFnInstance::new(Binary, Operator::Add).erased();
211        let scalar_fn_array = ScalarFnArray::try_new(scalar_fn, vec![lhs, rhs])?;
212
213        let result = scalar_fn_array
214            .into_array()
215            .execute::<Canonical>(&mut array_session().create_execution_ctx())?
216            .into_array();
217        let expected = PrimitiveArray::new(
218            buffer![11i32, 0, 33],
219            Validity::from_iter([true, false, true]),
220        )
221        .into_array();
222        assert_arrays_eq!(result, expected, &mut ctx);
223
224        Ok(())
225    }
226
227    #[test]
228    fn test_scalar_fn_comparison() -> VortexResult<()> {
229        let mut ctx = array_session().create_execution_ctx();
230        let lhs = buffer![1i32, 5, 3].into_array();
231        let rhs = buffer![2i32, 5, 1].into_array();
232
233        let scalar_fn = TypedScalarFnInstance::new(Binary, Operator::Eq).erased();
234        let scalar_fn_array = ScalarFnArray::try_new(scalar_fn, vec![lhs, rhs])?;
235
236        let result = scalar_fn_array
237            .into_array()
238            .execute::<Canonical>(&mut array_session().create_execution_ctx())?
239            .into_array();
240        let expected = BoolArray::from_iter([false, true, false]).into_array();
241        assert_arrays_eq!(result, expected, &mut ctx);
242
243        Ok(())
244    }
245}