Skip to main content

spark_connect/
datasource.rs

1//! Data source registration support.
2//!
3//! Mirrors `pyspark.sql.connect.datasource` and allows registration of custom
4//! Python data sources that can be used in SQL and DataFrame queries.
5//! Data sources are cloudpickled on the Python client and wrapped into
6//! `CommonInlineUserDefinedDataSource` commands for transmission to the server.
7
8use spark_connect_proto as proto;
9
10/// Represents a Python data source with its serialized command and metadata.
11#[derive(Debug, Clone, PartialEq)]
12pub struct PythonDataSourcePayload {
13    /// The cloudpickled command bytes containing the serialized data source.
14    pub command: Vec<u8>,
15    /// Python version used for pickling (e.g., "3.9", "3.11").
16    pub python_ver: String,
17}
18
19impl PythonDataSourcePayload {
20    /// Create a new Python data source payload.
21    pub fn new(command: Vec<u8>, python_ver: String) -> Self {
22        PythonDataSourcePayload {
23            command,
24            python_ver,
25        }
26    }
27
28    /// Convert to a proto PythonDataSource message.
29    pub fn to_proto(&self) -> proto::PythonDataSource {
30        use bytes::Bytes;
31        let mut proto = proto::PythonDataSource::default();
32        proto.command = Bytes::copy_from_slice(&self.command);
33        proto.python_ver = self.python_ver.clone();
34        proto
35    }
36}
37
38/// Represents a CommonInlineUserDefinedDataSource command.
39/// Wraps a Python data source with its name.
40#[derive(Debug, Clone, PartialEq)]
41pub struct CommonInlineUserDefinedDataSourceExpression {
42    /// Name of the data source (e.g., "my_source").
43    pub name: String,
44    /// The Python data source payload (command and python version).
45    pub python_data_source: PythonDataSourcePayload,
46}
47
48impl CommonInlineUserDefinedDataSourceExpression {
49    /// Create a new CommonInlineUserDefinedDataSource expression.
50    pub fn new(name: String, python_data_source: PythonDataSourcePayload) -> Self {
51        CommonInlineUserDefinedDataSourceExpression {
52            name,
53            python_data_source,
54        }
55    }
56
57    /// Convert to a proto CommonInlineUserDefinedDataSource message.
58    pub fn to_proto(&self) -> proto::CommonInlineUserDefinedDataSource {
59        let mut proto = proto::CommonInlineUserDefinedDataSource::default();
60        proto.name = self.name.clone();
61        proto.data_source = Some(
62            proto::common_inline_user_defined_data_source::DataSource::PythonDataSource(
63                self.python_data_source.to_proto(),
64            ),
65        );
66        proto
67    }
68}
69
70#[cfg(test)]
71mod tests {
72    use super::*;
73
74    #[test]
75    fn test_python_data_source_payload_to_proto() {
76        let payload = PythonDataSourcePayload::new(vec![1, 2, 3, 4, 5], "3.9".to_string());
77
78        let proto = payload.to_proto();
79
80        assert_eq!(proto.python_ver, "3.9");
81        assert_eq!(proto.command.len(), 5);
82    }
83
84    #[test]
85    fn test_common_inline_data_source_to_proto() {
86        let payload =
87            PythonDataSourcePayload::new(b"pickled_datasource_bytes".to_vec(), "3.11".to_string());
88
89        let ds_expr =
90            CommonInlineUserDefinedDataSourceExpression::new("my_source".to_string(), payload);
91
92        let proto = ds_expr.to_proto();
93
94        assert_eq!(proto.name, "my_source");
95        assert!(proto.data_source.is_some());
96
97        // Verify the Python data source is embedded
98        if let Some(proto::common_inline_user_defined_data_source::DataSource::PythonDataSource(
99            py_ds,
100        )) = proto.data_source
101        {
102            assert_eq!(py_ds.python_ver, "3.11");
103            assert_eq!(
104                py_ds.command,
105                bytes::Bytes::copy_from_slice(b"pickled_datasource_bytes")
106            );
107        } else {
108            panic!("Expected Python data source in CommonInlineUserDefinedDataSource");
109        }
110    }
111}