1use 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
41pub 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
129pub 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
141pub 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
154pub 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 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 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 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 buf.push(use_cloudpickle as u8);
247 buf.extend_from_slice(&*PYTHON3_VERSION);
248
249 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 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 #[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}