1use crate::expression::Expression;
8use crate::types::DataType;
9use spark_connect_proto as proto;
10
11pub mod eval_type {
14 pub const SQL_BATCHED_UDF: i32 = 100;
16 pub const SQL_ARROW_BATCHED_UDF: i32 = 101;
18
19 pub const SQL_SCALAR_PANDAS_UDF: i32 = 200;
21 pub const SQL_GROUPED_MAP_PANDAS_UDF: i32 = 201;
23 pub const SQL_GROUPED_AGG_PANDAS_UDF: i32 = 202;
25 pub const SQL_WINDOW_AGG_PANDAS_UDF: i32 = 203;
27 pub const SQL_SCALAR_PANDAS_ITER_UDF: i32 = 204;
29 pub const SQL_MAP_PANDAS_ITER_UDF: i32 = 205;
31 pub const SQL_COGROUPED_MAP_PANDAS_UDF: i32 = 206;
33 pub const SQL_MAP_ARROW_ITER_UDF: i32 = 207;
35 pub const SQL_GROUPED_MAP_PANDAS_UDF_WITH_STATE: i32 = 208;
37 pub const SQL_GROUPED_MAP_ARROW_UDF: i32 = 209;
39 pub const SQL_COGROUPED_MAP_ARROW_UDF: i32 = 210;
41 pub const SQL_TRANSFORM_WITH_STATE_PANDAS_UDF: i32 = 211;
43 pub const SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF: i32 = 212;
45 pub const SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF: i32 = 213;
47 pub const SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF: i32 = 214;
49 pub const SQL_GROUPED_MAP_ARROW_ITER_UDF: i32 = 215;
51 pub const SQL_GROUPED_MAP_PANDAS_ITER_UDF: i32 = 216;
53 pub const SQL_GROUPED_AGG_PANDAS_ITER_UDF: i32 = 217;
55
56 pub const SQL_SCALAR_ARROW_UDF: i32 = 250;
58 pub const SQL_SCALAR_ARROW_ITER_UDF: i32 = 251;
60 pub const SQL_GROUPED_AGG_ARROW_UDF: i32 = 252;
62 pub const SQL_WINDOW_AGG_ARROW_UDF: i32 = 253;
64 pub const SQL_GROUPED_AGG_ARROW_ITER_UDF: i32 = 254;
66
67 pub const SQL_TABLE_UDF: i32 = 300;
69 pub const SQL_ARROW_TABLE_UDF: i32 = 301;
71 pub const SQL_ARROW_UDTF: i32 = 302;
73}
74
75#[derive(Debug, Clone, PartialEq)]
77pub struct PythonUDFPayload {
78 pub output_type: DataType,
80 pub eval_type: i32,
82 pub command: Vec<u8>,
84 pub python_ver: String,
86}
87
88impl PythonUDFPayload {
89 pub fn new(
91 output_type: DataType,
92 eval_type: i32,
93 command: Vec<u8>,
94 python_ver: String,
95 ) -> Self {
96 PythonUDFPayload {
97 output_type,
98 eval_type,
99 command,
100 python_ver,
101 }
102 }
103
104 pub fn to_proto(&self) -> proto::PythonUdf {
106 use bytes::Bytes;
107 let mut proto = proto::PythonUdf::default();
108 proto.output_type = Some(self.output_type.to_proto());
109 proto.eval_type = self.eval_type;
110 proto.command = Bytes::copy_from_slice(&self.command);
111 proto.python_ver = self.python_ver.clone();
112 proto
113 }
114}
115
116#[derive(Debug, Clone, PartialEq)]
119pub struct CommonInlineUserDefinedFunctionExpression {
120 pub function_name: String,
122 pub deterministic: bool,
124 pub arguments: Vec<Expression>,
126 pub python_udf: PythonUDFPayload,
128}
129
130impl CommonInlineUserDefinedFunctionExpression {
131 pub fn new(
133 function_name: String,
134 deterministic: bool,
135 arguments: Vec<Expression>,
136 python_udf: PythonUDFPayload,
137 ) -> Self {
138 CommonInlineUserDefinedFunctionExpression {
139 function_name,
140 deterministic,
141 arguments,
142 python_udf,
143 }
144 }
145
146 pub fn to_proto(&self) -> proto::CommonInlineUserDefinedFunction {
148 let mut proto = proto::CommonInlineUserDefinedFunction::default();
149 proto.function_name = self.function_name.clone();
150 proto.deterministic = self.deterministic;
151 proto.arguments = self.arguments.iter().map(|expr| expr.to_proto()).collect();
152 proto.is_distinct = false;
153 proto.function = Some(
154 proto::common_inline_user_defined_function::Function::PythonUdf(
155 self.python_udf.to_proto(),
156 ),
157 );
158 proto
159 }
160}
161
162#[cfg(test)]
163mod tests {
164 use super::*;
165
166 #[test]
167 fn test_python_udf_payload_to_proto() {
168 let payload = PythonUDFPayload::new(
169 DataType::Integer,
170 eval_type::SQL_BATCHED_UDF,
171 vec![1, 2, 3, 4, 5],
172 "3.9".to_string(),
173 );
174
175 let proto = payload.to_proto();
176
177 assert_eq!(proto.eval_type, 100); assert_eq!(proto.python_ver, "3.9");
179 assert_eq!(proto.command.len(), 5);
180 assert!(proto.output_type.is_some());
181 }
182
183 #[test]
184 fn test_common_inline_udf_expression_to_proto() {
185 let payload = PythonUDFPayload::new(
186 DataType::String {
187 collation: "UTF8_BINARY".to_string(),
188 },
189 eval_type::SQL_SCALAR_PANDAS_UDF,
190 b"pickled_command_bytes".to_vec(),
191 "3.11".to_string(),
192 );
193
194 let udf_expr = CommonInlineUserDefinedFunctionExpression::new(
195 "my_udf".to_string(),
196 true,
197 vec![],
198 payload,
199 );
200
201 let proto = udf_expr.to_proto();
202
203 assert_eq!(proto.function_name, "my_udf");
204 assert_eq!(proto.deterministic, true);
205 assert_eq!(proto.is_distinct, false);
206 assert_eq!(proto.arguments.len(), 0);
207 assert!(proto.function.is_some());
208
209 if let Some(proto::common_inline_user_defined_function::Function::PythonUdf(py_udf)) =
211 proto.function
212 {
213 assert_eq!(py_udf.eval_type, 200); assert_eq!(py_udf.python_ver, "3.11");
215 assert_eq!(
216 py_udf.command,
217 bytes::Bytes::copy_from_slice(b"pickled_command_bytes")
218 );
219 } else {
220 panic!("Expected Python UDF in CommonInlineUserDefinedFunction");
221 }
222 }
223
224 #[test]
225 fn test_eval_type_constants() {
226 assert_eq!(eval_type::SQL_BATCHED_UDF, 100);
228 assert_eq!(eval_type::SQL_ARROW_BATCHED_UDF, 101);
229 assert_eq!(eval_type::SQL_SCALAR_PANDAS_UDF, 200);
230 assert_eq!(eval_type::SQL_GROUPED_MAP_PANDAS_UDF, 201);
231 assert_eq!(eval_type::SQL_GROUPED_AGG_PANDAS_UDF, 202);
232 assert_eq!(eval_type::SQL_WINDOW_AGG_PANDAS_UDF, 203);
233 assert_eq!(eval_type::SQL_SCALAR_PANDAS_ITER_UDF, 204);
234 assert_eq!(eval_type::SQL_MAP_PANDAS_ITER_UDF, 205);
235 assert_eq!(eval_type::SQL_COGROUPED_MAP_PANDAS_UDF, 206);
236 assert_eq!(eval_type::SQL_MAP_ARROW_ITER_UDF, 207);
237 assert_eq!(eval_type::SQL_GROUPED_MAP_ARROW_UDF, 209);
238 assert_eq!(eval_type::SQL_COGROUPED_MAP_ARROW_UDF, 210);
239 assert_eq!(eval_type::SQL_SCALAR_ARROW_UDF, 250);
240 assert_eq!(eval_type::SQL_SCALAR_ARROW_ITER_UDF, 251);
241 assert_eq!(eval_type::SQL_GROUPED_AGG_ARROW_UDF, 252);
242 assert_eq!(eval_type::SQL_WINDOW_AGG_ARROW_UDF, 253);
243 assert_eq!(eval_type::SQL_GROUPED_AGG_ARROW_ITER_UDF, 254);
244 assert_eq!(eval_type::SQL_TABLE_UDF, 300);
245 assert_eq!(eval_type::SQL_ARROW_TABLE_UDF, 301);
246 assert_eq!(eval_type::SQL_ARROW_UDTF, 302);
247 }
248}