common_function/scalars/
udf.rs1use 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
82pub 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 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}