Skip to main content

common_function/scalars/
udf.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::{Debug, Formatter};
16use std::hash::{Hash, Hasher};
17
18use datafusion::arrow::datatypes::DataType;
19use datafusion::logical_expr::{ScalarFunctionArgs, ScalarUDFImpl};
20use datafusion_expr::ScalarUDF;
21
22use crate::function::FunctionRef;
23
24struct ScalarUdf {
25    function: FunctionRef,
26}
27
28impl Debug for ScalarUdf {
29    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
30        f.debug_struct("ScalarUdf")
31            .field("function", &self.function.name())
32            .finish()
33    }
34}
35
36impl PartialEq for ScalarUdf {
37    fn eq(&self, other: &Self) -> bool {
38        self.function.signature() == other.function.signature()
39    }
40}
41
42impl Eq for ScalarUdf {}
43
44impl Hash for ScalarUdf {
45    fn hash<H: Hasher>(&self, state: &mut H) {
46        self.function.signature().hash(state)
47    }
48}
49
50impl ScalarUDFImpl for ScalarUdf {
51    fn name(&self) -> &str {
52        self.function.name()
53    }
54
55    fn aliases(&self) -> &[String] {
56        self.function.aliases()
57    }
58
59    fn signature(&self) -> &datafusion_expr::Signature {
60        self.function.signature()
61    }
62
63    fn return_type(&self, arg_types: &[DataType]) -> datafusion_common::Result<DataType> {
64        self.function.return_type(arg_types)
65    }
66
67    fn return_field_from_args(
68        &self,
69        args: datafusion_expr::ReturnFieldArgs,
70    ) -> datafusion_common::Result<arrow_schema::FieldRef> {
71        self.function.return_field_from_args(args)
72    }
73
74    fn invoke_with_args(
75        &self,
76        args: ScalarFunctionArgs,
77    ) -> datafusion_common::Result<datafusion_expr::ColumnarValue> {
78        self.function.invoke_with_args(args)
79    }
80}
81
82/// Create a ScalarUdf from function, query context and state.
83pub fn create_udf(function: FunctionRef) -> ScalarUDF {
84    ScalarUDF::new_from_impl(ScalarUdf { function })
85}
86
87#[cfg(test)]
88mod tests {
89    use std::sync::Arc;
90
91    use common_query::prelude::ScalarValue;
92    use datafusion::arrow::array::BooleanArray;
93    use datafusion_common::arrow::array::AsArray;
94    use datafusion_common::arrow::datatypes::DataType as ArrowDataType;
95    use datafusion_common::config::ConfigOptions;
96    use datatypes::arrow::datatypes::Field;
97    use datatypes::data_type::{ConcreteDataType, DataType};
98
99    use super::*;
100    use crate::function::Function;
101    use crate::scalars::test::TestAndFunction;
102
103    #[test]
104    fn test_create_udf() {
105        let f = Arc::new(TestAndFunction::default());
106
107        let args = ScalarFunctionArgs {
108            args: vec![
109                datafusion_expr::ColumnarValue::Array(Arc::new(BooleanArray::from(vec![
110                    true, true, true,
111                ]))),
112                datafusion_expr::ColumnarValue::Array(Arc::new(BooleanArray::from(vec![
113                    true, false, true,
114                ]))),
115            ],
116            arg_fields: vec![],
117            number_rows: 3,
118            return_field: Arc::new(Field::new("x", ArrowDataType::Boolean, true)),
119            config_options: Arc::new(Default::default()),
120        };
121
122        let result = f
123            .invoke_with_args(args)
124            .and_then(|x| x.to_array(3))
125            .unwrap();
126        let vector = result.as_boolean();
127        assert_eq!(3, vector.len());
128
129        assert!(vector.value(0));
130        assert!(!vector.value(1));
131        assert!(vector.value(2));
132
133        // create a udf and test it again
134        let udf = create_udf(f);
135
136        assert_eq!("test_and", udf.name());
137        assert_eq!(
138            ConcreteDataType::boolean_datatype(),
139            udf.return_type(&[])
140                .map(|x| ConcreteDataType::from_arrow_type(&x))
141                .unwrap()
142        );
143
144        let args = vec![
145            datafusion_expr::ColumnarValue::Scalar(ScalarValue::Boolean(Some(true))),
146            datafusion_expr::ColumnarValue::Array(Arc::new(BooleanArray::from(vec![
147                true, false, false, true,
148            ]))),
149        ];
150
151        let arg_fields = vec![
152            Arc::new(Field::new("a", args[0].data_type(), false)),
153            Arc::new(Field::new("b", args[1].data_type(), false)),
154        ];
155        let return_field = Arc::new(Field::new(
156            "x",
157            ConcreteDataType::boolean_datatype().as_arrow_type(),
158            false,
159        ));
160        let args = ScalarFunctionArgs {
161            args,
162            arg_fields,
163            number_rows: 4,
164            return_field,
165            config_options: Arc::new(ConfigOptions::default()),
166        };
167        match udf.invoke_with_args(args).unwrap() {
168            datafusion_expr::ColumnarValue::Array(x) => {
169                let x = x.as_any().downcast_ref::<BooleanArray>().unwrap();
170                assert_eq!(x.len(), 4);
171                assert_eq!(
172                    x.iter().flatten().collect::<Vec<bool>>(),
173                    vec![true, false, false, true]
174                );
175            }
176            _ => unreachable!(),
177        }
178    }
179}