Skip to main content

spark_connect/
udf.rs

1//! User-defined function (UDF) support.
2//!
3//! Mirrors `pyspark.sql.connect.udf` and `pyspark.util.PythonEvalType`.
4//! UDFs are cloudpickled on the Python client and wrapped into
5//! `CommonInlineUserDefinedFunction` expressions for transmission to the server.
6
7use crate::expression::Expression;
8use crate::types::DataType;
9use spark_connect_proto as proto;
10
11/// Python evaluation type constants, matching `pyspark.util.PythonEvalType`.
12/// These distinguish between different UDF execution modes (batched, pandas, arrow, etc.).
13pub mod eval_type {
14    /// Regular Python UDF, row-by-row (column as list).
15    pub const SQL_BATCHED_UDF: i32 = 100;
16    /// Arrow-optimized Python UDF (column as PyArrow table).
17    pub const SQL_ARROW_BATCHED_UDF: i32 = 101;
18
19    /// Pandas scalar UDF (Series -> Series).
20    pub const SQL_SCALAR_PANDAS_UDF: i32 = 200;
21    /// Pandas grouped map UDF (grouped DataFrame -> DataFrame).
22    pub const SQL_GROUPED_MAP_PANDAS_UDF: i32 = 201;
23    /// Pandas grouped aggregate UDF.
24    pub const SQL_GROUPED_AGG_PANDAS_UDF: i32 = 202;
25    /// Pandas window aggregate UDF.
26    pub const SQL_WINDOW_AGG_PANDAS_UDF: i32 = 203;
27    /// Pandas scalar iterator UDF.
28    pub const SQL_SCALAR_PANDAS_ITER_UDF: i32 = 204;
29    /// Pandas map iterator UDF.
30    pub const SQL_MAP_PANDAS_ITER_UDF: i32 = 205;
31    /// Pandas cogrouped map UDF.
32    pub const SQL_COGROUPED_MAP_PANDAS_UDF: i32 = 206;
33    /// Arrow map iterator UDF.
34    pub const SQL_MAP_ARROW_ITER_UDF: i32 = 207;
35    /// Pandas grouped map with state UDF.
36    pub const SQL_GROUPED_MAP_PANDAS_UDF_WITH_STATE: i32 = 208;
37    /// Arrow grouped map UDF.
38    pub const SQL_GROUPED_MAP_ARROW_UDF: i32 = 209;
39    /// Arrow cogrouped map UDF.
40    pub const SQL_COGROUPED_MAP_ARROW_UDF: i32 = 210;
41    /// Pandas transform with state UDF.
42    pub const SQL_TRANSFORM_WITH_STATE_PANDAS_UDF: i32 = 211;
43    /// Pandas transform with state init state UDF.
44    pub const SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF: i32 = 212;
45    /// Python row transform with state UDF.
46    pub const SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF: i32 = 213;
47    /// Python row transform with state init state UDF.
48    pub const SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF: i32 = 214;
49    /// Arrow grouped map iterator UDF.
50    pub const SQL_GROUPED_MAP_ARROW_ITER_UDF: i32 = 215;
51    /// Pandas grouped map iterator UDF.
52    pub const SQL_GROUPED_MAP_PANDAS_ITER_UDF: i32 = 216;
53    /// Pandas grouped aggregate iterator UDF.
54    pub const SQL_GROUPED_AGG_PANDAS_ITER_UDF: i32 = 217;
55
56    /// Arrow scalar UDF.
57    pub const SQL_SCALAR_ARROW_UDF: i32 = 250;
58    /// Arrow scalar iterator UDF.
59    pub const SQL_SCALAR_ARROW_ITER_UDF: i32 = 251;
60    /// Arrow grouped aggregate UDF.
61    pub const SQL_GROUPED_AGG_ARROW_UDF: i32 = 252;
62    /// Arrow window aggregate UDF.
63    pub const SQL_WINDOW_AGG_ARROW_UDF: i32 = 253;
64    /// Arrow grouped aggregate iterator UDF.
65    pub const SQL_GROUPED_AGG_ARROW_ITER_UDF: i32 = 254;
66
67    /// SQL table UDF (UDTF).
68    pub const SQL_TABLE_UDF: i32 = 300;
69    /// Arrow SQL table UDF (UDTF).
70    pub const SQL_ARROW_TABLE_UDF: i32 = 301;
71    /// Arrow UDTF.
72    pub const SQL_ARROW_UDTF: i32 = 302;
73}
74
75/// Represents a Python UDF with its serialized command and metadata.
76#[derive(Debug, Clone, PartialEq)]
77pub struct PythonUDFPayload {
78    /// The output data type of the UDF.
79    pub output_type: DataType,
80    /// The evaluation type (e.g., SQL_BATCHED_UDF, SQL_SCALAR_PANDAS_UDF).
81    pub eval_type: i32,
82    /// The cloudpickled command bytes: typically cloudpickle.dumps((func, output_type)).
83    pub command: Vec<u8>,
84    /// Python version used for pickling (e.g., "3.9", "3.11").
85    pub python_ver: String,
86}
87
88impl PythonUDFPayload {
89    /// Create a new Python UDF payload.
90    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    /// Convert to a proto PythonUDF message.
105    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/// Represents a CommonInlineUserDefinedFunction expression.
117/// Wraps a Python UDF with its name, determinism flag, and arguments.
118#[derive(Debug, Clone, PartialEq)]
119pub struct CommonInlineUserDefinedFunctionExpression {
120    /// Name of the UDF (e.g., "my_func").
121    pub function_name: String,
122    /// Whether the UDF is deterministic.
123    pub deterministic: bool,
124    /// Argument expressions passed to the UDF.
125    pub arguments: Vec<Expression>,
126    /// The Python UDF payload (command, output type, eval type, python version).
127    pub python_udf: PythonUDFPayload,
128}
129
130impl CommonInlineUserDefinedFunctionExpression {
131    /// Create a new CommonInlineUserDefinedFunction expression.
132    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    /// Convert to a proto CommonInlineUserDefinedFunction message.
147    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); // SQL_BATCHED_UDF
178        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        // Verify the Python UDF is embedded
210        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); // SQL_SCALAR_PANDAS_UDF
214            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        // Verify eval type constants match pyspark.util.PythonEvalType
227        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}