Skip to main content

common_function/scalars/math/
modulo.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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/// The function to find remainders
29#[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}