Skip to main content

polars_utils/
pl_serialize.rs

1//! Centralized Polars serialization entry.
2//!
3//! Currently provides two serialization scheme's.
4//! - Self-describing (and thus more forward compatible) activated with `FC: true`
5//! - Compact activated with `FC: false`
6use std::io::{Read, Write};
7
8use polars_error::{PolarsResult, to_compute_err};
9
10fn config() -> bincode::config::Configuration {
11    bincode::config::standard()
12        .with_no_limit()
13        .with_variable_int_encoding()
14}
15
16fn serialize_impl<T, const FC: bool>(mut writer: &mut dyn Write, value: &T) -> PolarsResult<()>
17where
18    T: serde::ser::Serialize,
19{
20    if FC {
21        let mut s = rmp_serde::Serializer::new(&mut *writer).with_struct_map();
22        value.serialize(&mut s).map_err(to_compute_err)
23    } else {
24        bincode::serde::encode_into_std_write(value, &mut writer, config())
25            .map_err(to_compute_err)
26            .map(|_| ())
27    }
28}
29
30fn deserialize_impl<T, const FC: bool>(mut reader: &mut dyn Read) -> PolarsResult<T>
31where
32    T: serde::de::DeserializeOwned,
33{
34    if FC {
35        rmp_serde::from_read(&mut *reader).map_err(to_compute_err)
36    } else {
37        bincode::serde::decode_from_std_read(&mut reader, config()).map_err(to_compute_err)
38    }
39}
40
41/// Mainly used to enable compression when serializing the final outer value.
42/// For intermediate serialization steps, the function in the module should
43/// be used instead.
44pub struct SerializeOptions {
45    compression: bool,
46}
47
48impl SerializeOptions {
49    pub fn with_compression(mut self, compression: bool) -> Self {
50        self.compression = compression;
51        self
52    }
53
54    pub fn serialize_into_writer<W, T, const FC: bool>(
55        &self,
56        mut writer: W,
57        value: &T,
58    ) -> PolarsResult<()>
59    where
60        W: std::io::Write,
61        T: serde::ser::Serialize,
62    {
63        if self.compression {
64            let mut compr_writer =
65                flate2::write::ZlibEncoder::new(writer, flate2::Compression::fast());
66            serialize_impl::<_, FC>(&mut compr_writer, value)
67        } else {
68            serialize_impl::<_, FC>(&mut writer, value)
69        }
70    }
71
72    pub fn deserialize_from_reader<T, R, const FC: bool>(&self, mut reader: R) -> PolarsResult<T>
73    where
74        T: serde::de::DeserializeOwned,
75        R: std::io::Read,
76    {
77        if self.compression {
78            let mut compr_reader = flate2::read::ZlibDecoder::new(reader);
79            deserialize_impl::<_, FC>(&mut compr_reader)
80        } else {
81            deserialize_impl::<_, FC>(&mut reader)
82        }
83    }
84
85    pub fn serialize_to_bytes<T, const FC: bool>(&self, value: &T) -> PolarsResult<Vec<u8>>
86    where
87        T: serde::ser::Serialize,
88    {
89        let mut v = vec![];
90
91        self.serialize_into_writer::<_, _, FC>(&mut v, value)?;
92
93        Ok(v)
94    }
95}
96
97#[allow(clippy::derivable_impls)]
98impl Default for SerializeOptions {
99    fn default() -> Self {
100        Self { compression: false }
101    }
102}
103
104pub fn serialize_into_writer<W, T, const FC: bool>(mut writer: W, value: &T) -> PolarsResult<()>
105where
106    W: std::io::Write,
107    T: serde::ser::Serialize,
108{
109    serialize_impl::<_, FC>(&mut writer, value)
110}
111
112pub fn deserialize_from_reader<T, R, const FC: bool>(mut reader: R) -> PolarsResult<T>
113where
114    T: serde::de::DeserializeOwned,
115    R: std::io::Read,
116{
117    deserialize_impl::<_, FC>(&mut reader)
118}
119
120pub fn serialize_to_bytes<T, const FC: bool>(value: &T) -> PolarsResult<Vec<u8>>
121where
122    T: serde::ser::Serialize,
123{
124    let mut v = vec![];
125    serialize_into_writer::<_, _, FC>(&mut v, value)?;
126    Ok(v)
127}
128
129/// Serialize function customized for `DslPlan`, with stack overflow protection.
130pub fn serialize_dsl<W, T>(mut writer: W, value: &T) -> PolarsResult<()>
131where
132    W: std::io::Write,
133    T: serde::ser::Serialize,
134{
135    let writer: &mut dyn std::io::Write = &mut writer;
136    let mut s = rmp_serde::Serializer::new(writer).with_struct_map();
137    let s = serde_stacker::Serializer::new(&mut s);
138    value.serialize(s).map_err(to_compute_err)
139}
140
141/// Deserialize function customized for `DslPlan`, with stack overflow protection.
142pub fn deserialize_dsl<T, R>(mut reader: R) -> PolarsResult<T>
143where
144    T: serde::de::DeserializeOwned,
145    R: std::io::Read,
146{
147    let reader: &mut dyn std::io::Read = &mut reader;
148    let mut de = rmp_serde::Deserializer::new(reader);
149    de.set_max_depth(usize::MAX);
150    let de = serde_stacker::Deserializer::new(&mut de);
151    T::deserialize(de).map_err(to_compute_err)
152}
153
154/// Potentially avoids copying memory compared to a naive `Vec::<u8>::deserialize`.
155///
156/// This is essentially boilerplate for visiting bytes without copying where possible.
157pub fn deserialize_map_bytes<'de, D, O>(
158    deserializer: D,
159    mut func: impl for<'b> FnMut(std::borrow::Cow<'b, [u8]>) -> O,
160) -> Result<O, D::Error>
161where
162    D: serde::de::Deserializer<'de>,
163{
164    // Lets us avoid monomorphizing the visitor
165    let mut out: Option<O> = None;
166    struct V<'f>(&'f mut dyn for<'b> FnMut(std::borrow::Cow<'b, [u8]>));
167
168    deserializer.deserialize_bytes(V(&mut |v| drop(out.replace(func(v)))))?;
169
170    return Ok(out.unwrap());
171
172    impl<'de> serde::de::Visitor<'de> for V<'_> {
173        type Value = ();
174
175        fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
176            formatter.write_str("deserialize_map_bytes")
177        }
178
179        fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
180        where
181            E: serde::de::Error,
182        {
183            self.0(std::borrow::Cow::Borrowed(v));
184            Ok(())
185        }
186
187        fn visit_byte_buf<E>(self, v: Vec<u8>) -> Result<Self::Value, E>
188        where
189            E: serde::de::Error,
190        {
191            self.0(std::borrow::Cow::Owned(v));
192            Ok(())
193        }
194
195        fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
196        where
197            A: serde::de::SeqAccess<'de>,
198        {
199            // This is not ideal, but we hit here if the serialization format is JSON.
200            let bytes = std::iter::from_fn(|| seq.next_element::<u8>().transpose())
201                .collect::<Result<Vec<_>, A::Error>>()?;
202
203            self.0(std::borrow::Cow::Owned(bytes));
204            Ok(())
205        }
206    }
207}
208
209thread_local! {
210    pub static USE_CLOUDPICKLE: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
211}
212
213#[cfg(feature = "python")]
214pub fn python_object_serialize(
215    pyobj: &pyo3::Py<pyo3::PyAny>,
216    buf: &mut Vec<u8>,
217) -> PolarsResult<()> {
218    use pyo3::Python;
219    use pyo3::pybacked::PyBackedBytes;
220    use pyo3::types::{PyAnyMethods, PyModule};
221
222    use crate::python_function::PYTHON3_VERSION;
223
224    let mut use_cloudpickle = USE_CLOUDPICKLE.get();
225    let dumped = Python::attach(|py| {
226        // Pickle with whatever pickling method was selected.
227        if use_cloudpickle {
228            let cloudpickle = PyModule::import(py, "cloudpickle")?.getattr("dumps")?;
229            cloudpickle.call1((pyobj.clone_ref(py),))?
230        } else {
231            let pickle = PyModule::import(py, "pickle")?.getattr("dumps")?;
232            match pickle.call1((pyobj.clone_ref(py),)) {
233                Ok(dumped) => dumped,
234                Err(_) => {
235                    use_cloudpickle = true;
236                    let cloudpickle = PyModule::import(py, "cloudpickle")?.getattr("dumps")?;
237                    cloudpickle.call1((pyobj.clone_ref(py),))?
238                },
239            }
240        }
241        .extract::<PyBackedBytes>()
242        .map_err(pyo3::PyErr::from)
243    })?;
244
245    // Write pickle metadata
246    buf.push(use_cloudpickle as u8);
247    buf.extend_from_slice(&*PYTHON3_VERSION);
248
249    // Write UDF
250    buf.extend_from_slice(&dumped);
251    Ok(())
252}
253
254#[cfg(feature = "python")]
255pub fn python_object_deserialize(buf: &[u8]) -> PolarsResult<pyo3::Py<pyo3::PyAny>> {
256    use polars_error::polars_ensure;
257    use pyo3::Python;
258    use pyo3::types::{PyAnyMethods, PyBytes, PyModule};
259
260    use crate::python_function::PYTHON3_VERSION;
261
262    // Handle pickle metadata
263    let use_cloudpickle = buf[0] != 0;
264    if use_cloudpickle {
265        let ser_py_version = &buf[1..3];
266        let cur_py_version = *PYTHON3_VERSION;
267        polars_ensure!(
268            ser_py_version == cur_py_version,
269            InvalidOperation:
270            "current Python version {:?} does not match the Python version used to serialize the UDF {:?}",
271            (3, cur_py_version[0], cur_py_version[1]),
272            (3, ser_py_version[0], ser_py_version[1] )
273        );
274    }
275    let buf = &buf[3..];
276
277    Python::attach(|py| {
278        let loads = PyModule::import(py, "pickle")?.getattr("loads")?;
279        let arg = (PyBytes::new(py, buf),);
280        let python_function = loads.call1(arg)?;
281        Ok(python_function.into())
282    })
283}
284
285#[cfg(test)]
286mod tests {
287    #[test]
288    fn test_serde_skip_enum() {
289        #[derive(Default, Debug, PartialEq)]
290        struct MyType(Option<usize>);
291
292        // Note: serde(skip) must be at the end of enums
293        #[derive(Debug, PartialEq, serde::Serialize, serde::Deserialize)]
294        enum Enum {
295            A,
296            #[serde(skip)]
297            B(MyType),
298        }
299
300        impl Default for Enum {
301            fn default() -> Self {
302                Self::B(MyType(None))
303            }
304        }
305
306        let v = Enum::A;
307        let b = super::serialize_to_bytes::<_, false>(&v).unwrap();
308        let r: Enum = super::deserialize_from_reader::<_, _, false>(b.as_slice()).unwrap();
309
310        assert_eq!(r, v);
311
312        let v = Enum::A;
313        let b = super::SerializeOptions::default()
314            .serialize_to_bytes::<_, false>(&v)
315            .unwrap();
316        let r: Enum = super::SerializeOptions::default()
317            .deserialize_from_reader::<_, _, false>(b.as_slice())
318            .unwrap();
319
320        assert_eq!(r, v);
321    }
322}