vortex_array/arrays/scalar_fn/vtable/
operations.rs1use 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}