common_function/scalars/math/
modulo.rs1use std::fmt;
16use std::fmt::Display;
17
18use datafusion_common::arrow::compute;
19use datafusion_common::arrow::compute::kernels::numeric;
20use datafusion_common::arrow::datatypes::DataType;
21use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, Signature, Volatility};
22
23use crate::function::{Function, extract_args};
24use crate::helper::NUMERICS;
25
26const NAME: &str = "mod";
27
28#[derive(Clone, Debug)]
30pub(crate) struct ModuloFunction {
31 signature: Signature,
32}
33
34impl Default for ModuloFunction {
35 fn default() -> Self {
36 Self {
37 signature: Signature::uniform(2, NUMERICS.to_vec(), Volatility::Immutable),
38 }
39 }
40}
41
42impl Display for ModuloFunction {
43 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
44 write!(f, "{}", NAME.to_ascii_uppercase())
45 }
46}
47
48impl Function for ModuloFunction {
49 fn name(&self) -> &str {
50 NAME
51 }
52
53 fn return_type(&self, input_types: &[DataType]) -> datafusion_common::Result<DataType> {
54 if input_types.iter().all(DataType::is_signed_integer) {
55 Ok(DataType::Int64)
56 } else if input_types.iter().all(DataType::is_unsigned_integer) {
57 Ok(DataType::UInt64)
58 } else {
59 Ok(DataType::Float64)
60 }
61 }
62
63 fn signature(&self) -> &Signature {
64 &self.signature
65 }
66
67 fn invoke_with_args(
68 &self,
69 args: ScalarFunctionArgs,
70 ) -> datafusion_common::Result<ColumnarValue> {
71 let [nums, divs] = extract_args(self.name(), &args)?;
72 let array = numeric::rem(&nums, &divs)?;
73
74 let result = match nums.data_type() {
75 DataType::Int8 | DataType::Int16 | DataType::Int32 | DataType::Int64 => {
76 compute::cast(&array, &DataType::Int64)
77 }
78 DataType::UInt8 | DataType::UInt16 | DataType::UInt32 | DataType::UInt64 => {
79 compute::cast(&array, &DataType::UInt64)
80 }
81 DataType::Float32 | DataType::Float64 => compute::cast(&array, &DataType::Float64),
82 _ => unreachable!("unexpected datatype: {:?}", nums.data_type()),
83 }?;
84 Ok(ColumnarValue::Array(result))
85 }
86}
87
88#[cfg(test)]
89mod tests {
90 use std::sync::Arc;
91
92 use arrow_schema::Field;
93 use datafusion_common::ScalarValue;
94 use datafusion_common::arrow::array::{
95 AsArray, Decimal128Array, Float64Array, Int32Array, StringViewArray, UInt32Array,
96 };
97 use datafusion_common::arrow::datatypes::{Float64Type, Int64Type, UInt64Type};
98
99 use super::*;
100 fn decimal_array(values: Vec<i128>) -> ColumnarValue {
101 ColumnarValue::Array(Arc::new(
102 Decimal128Array::from(values)
103 .with_precision_and_scale(10, 2)
104 .unwrap(),
105 ))
106 }
107
108 fn decimal_scalar(value: i128) -> ColumnarValue {
109 ColumnarValue::Scalar(ScalarValue::Decimal128(Some(value), 10, 2))
110 }
111
112 #[test]
113 #[allow(deprecated)]
114 fn modulo_decimal_coercion_executes_as_float64() {
115 let function = ModuloFunction::default();
116 let test_cases = [
117 (
118 vec![decimal_array(vec![500, 600]), decimal_scalar(200)],
119 vec![1.0, 0.0],
120 ),
121 (
122 vec![decimal_scalar(500), decimal_array(vec![200, 300])],
123 vec![1.0, 2.0],
124 ),
125 ];
126
127 for (args, expected) in test_cases {
128 let input_types = args
129 .iter()
130 .map(ColumnarValue::data_type)
131 .collect::<Vec<_>>();
132 let planned_types = datafusion_expr::type_coercion::functions::data_types(
133 function.name(),
134 &input_types,
135 function.signature(),
136 )
137 .unwrap();
138 assert_eq!(vec![DataType::Float64; 2], planned_types);
139 let args = args
140 .into_iter()
141 .zip(planned_types)
142 .map(|(arg, planned_type)| arg.cast_to(&planned_type, None))
143 .collect::<datafusion_common::Result<Vec<_>>>()
144 .unwrap();
145 let result = function
146 .invoke_with_args(ScalarFunctionArgs {
147 args,
148 arg_fields: vec![],
149 number_rows: 2,
150 return_field: Arc::new(Field::new("x", DataType::Float64, false)),
151 config_options: Arc::new(Default::default()),
152 })
153 .unwrap()
154 .to_array(2)
155 .unwrap();
156 let result = result.as_primitive::<Float64Type>();
157 assert_eq!(&Float64Array::from(expected), result);
158 }
159 }
160
161 #[test]
162 fn test_mod_function_signed() {
163 let function = ModuloFunction::default();
164 assert_eq!("mod", function.name());
165 assert_eq!(
166 DataType::Int64,
167 function.return_type(&[DataType::Int64]).unwrap()
168 );
169 assert_eq!(
170 DataType::Int64,
171 function.return_type(&[DataType::Int32]).unwrap()
172 );
173
174 let nums = vec![18, -17, 5, -6];
175 let divs = vec![4, 8, -5, -5];
176
177 let args = ScalarFunctionArgs {
178 args: vec![
179 ColumnarValue::Array(Arc::new(Int32Array::from(nums.clone()))),
180 ColumnarValue::Array(Arc::new(Int32Array::from(divs.clone()))),
181 ],
182 arg_fields: vec![],
183 number_rows: 4,
184 return_field: Arc::new(Field::new("x", DataType::Int64, false)),
185 config_options: Arc::new(Default::default()),
186 };
187 let result = function.invoke_with_args(args).unwrap();
188 let result = result.to_array(4).unwrap();
189 let result = result.as_primitive::<Int64Type>();
190 assert_eq!(result.len(), 4);
191 for i in 0..4 {
192 let p: i64 = (nums[i] % divs[i]) as i64;
193 assert_eq!(result.value(i), p);
194 }
195 }
196
197 #[test]
198 fn test_mod_function_unsigned() {
199 let function = ModuloFunction::default();
200 assert_eq!("mod", function.name());
201 assert_eq!(
202 DataType::UInt64,
203 function.return_type(&[DataType::UInt64]).unwrap()
204 );
205 assert_eq!(
206 DataType::UInt64,
207 function.return_type(&[DataType::UInt32]).unwrap()
208 );
209
210 let nums: Vec<u32> = vec![18, 17, 5, 6];
211 let divs: Vec<u32> = vec![4, 8, 5, 5];
212
213 let args = ScalarFunctionArgs {
214 args: vec![
215 ColumnarValue::Array(Arc::new(UInt32Array::from(nums.clone()))),
216 ColumnarValue::Array(Arc::new(UInt32Array::from(divs.clone()))),
217 ],
218 arg_fields: vec![],
219 number_rows: 4,
220 return_field: Arc::new(Field::new("x", DataType::UInt64, false)),
221 config_options: Arc::new(Default::default()),
222 };
223 let result = function.invoke_with_args(args).unwrap();
224 let result = result.to_array(4).unwrap();
225 let result = result.as_primitive::<UInt64Type>();
226 assert_eq!(result.len(), 4);
227 for i in 0..4 {
228 let p: u64 = (nums[i] % divs[i]) as u64;
229 assert_eq!(result.value(i), p);
230 }
231 }
232
233 #[test]
234 fn test_mod_function_float() {
235 let function = ModuloFunction::default();
236 assert_eq!("mod", function.name());
237 assert_eq!(
238 DataType::Float64,
239 function.return_type(&[DataType::Float64]).unwrap()
240 );
241 assert_eq!(
242 DataType::Float64,
243 function.return_type(&[DataType::Float32]).unwrap()
244 );
245
246 let nums = vec![18.0, 17.0, 5.0, 6.0];
247 let divs = vec![4.0, 8.0, 5.0, 5.0];
248
249 let args = ScalarFunctionArgs {
250 args: vec![
251 ColumnarValue::Array(Arc::new(Float64Array::from(nums.clone()))),
252 ColumnarValue::Array(Arc::new(Float64Array::from(divs.clone()))),
253 ],
254 arg_fields: vec![],
255 number_rows: 4,
256 return_field: Arc::new(Field::new("x", DataType::Float64, false)),
257 config_options: Arc::new(Default::default()),
258 };
259 let result = function.invoke_with_args(args).unwrap();
260 let result = result.to_array(4).unwrap();
261 let result = result.as_primitive::<Float64Type>();
262 assert_eq!(result.len(), 4);
263 for i in 0..4 {
264 let p: f64 = nums[i] % divs[i];
265 assert_eq!(result.value(i), p);
266 }
267 }
268
269 #[test]
270 fn test_mod_function_errors() {
271 let function = ModuloFunction::default();
272 assert_eq!("mod", function.name());
273 let nums = vec![27];
274 let divs = vec![0];
275
276 let args = ScalarFunctionArgs {
277 args: vec![
278 ColumnarValue::Array(Arc::new(Int32Array::from(nums))),
279 ColumnarValue::Array(Arc::new(Int32Array::from(divs))),
280 ],
281 arg_fields: vec![],
282 number_rows: 1,
283 return_field: Arc::new(Field::new("x", DataType::Int64, false)),
284 config_options: Arc::new(Default::default()),
285 };
286 let result = function.invoke_with_args(args);
287 assert!(result.is_err());
288 let err_msg = result.unwrap_err().to_string();
289 assert_eq!(err_msg, "Arrow error: Divide by zero error");
290
291 let nums = vec![27];
292
293 let args = ScalarFunctionArgs {
294 args: vec![ColumnarValue::Array(Arc::new(Int32Array::from(nums)))],
295 arg_fields: vec![],
296 number_rows: 1,
297 return_field: Arc::new(Field::new("x", DataType::Int64, false)),
298 config_options: Arc::new(Default::default()),
299 };
300 let result = function.invoke_with_args(args);
301 assert!(result.is_err());
302 let err_msg = result.unwrap_err().to_string();
303 assert_eq!(
304 err_msg,
305 "Execution error: mod function requires 2 arguments, got 1"
306 );
307
308 let nums = vec!["27"];
309 let divs = vec!["4"];
310 let args = ScalarFunctionArgs {
311 args: vec![
312 ColumnarValue::Array(Arc::new(StringViewArray::from(nums))),
313 ColumnarValue::Array(Arc::new(StringViewArray::from(divs))),
314 ],
315 arg_fields: vec![],
316 number_rows: 1,
317 return_field: Arc::new(Field::new("x", DataType::Int64, false)),
318 config_options: Arc::new(Default::default()),
319 };
320 let result = function.invoke_with_args(args);
321 assert!(result.is_err());
322 let err_msg = result.unwrap_err().to_string();
323 assert!(err_msg.contains("Invalid arithmetic operation"));
324 }
325}