Skip to main content

rustpython_vm/stdlib/
marshal.rs

1// spell-checker:ignore pyfrozen pycomplex
2pub(crate) use decl::module_def;
3
4#[pymodule(name = "marshal")]
5mod decl {
6    use crate::builtins::code::{CodeObject, Literal, PyVmBag};
7    use crate::class::StaticType;
8    use crate::common::wtf8::Wtf8;
9    use crate::{
10        PyObject, PyObjectRef, PyResult, TryFromObject, VirtualMachine,
11        builtins::{
12            PyBaseExceptionRef, PyBool, PyByteArray, PyBytes, PyCode, PyComplex, PyDict,
13            PyEllipsis, PyFloat, PyFrozenSet, PyInt, PyList, PyMemoryView, PyNone, PySet,
14            PyStopIteration, PyStr, PyTuple,
15        },
16        convert::ToPyObject,
17        function::ArgBytesLike,
18        object::{AsObject, PyPayload},
19    };
20    use core::cell::RefCell;
21    use malachite_bigint::BigInt;
22    use num_traits::Zero;
23    use rustpython_compiler_core::marshal::{self, DumpableValue};
24
25    #[pyattr(name = "version")]
26    use marshal::FORMAT_VERSION;
27
28    pub struct DumpError;
29
30    impl marshal::Dumpable for PyObjectRef {
31        type Error = DumpError;
32        type Constant = Literal;
33
34        fn with_dump<R>(
35            &self,
36            f: impl FnOnce(DumpableValue<'_, Self>) -> R,
37        ) -> Result<R, Self::Error> {
38            if self.is(PyStopIteration::static_type()) {
39                return Ok(f(DumpableValue::StopIter));
40            }
41
42            let ret = match_class!(match self {
43                PyNone => f(DumpableValue::None),
44                PyEllipsis => f(DumpableValue::Ellipsis),
45                ref pyint @ PyInt => {
46                    if self.class().is(PyBool::static_type()) {
47                        f(DumpableValue::Boolean(!pyint.as_bigint().is_zero()))
48                    } else {
49                        f(DumpableValue::Integer(pyint.as_bigint()))
50                    }
51                }
52                ref pyfloat @ PyFloat => {
53                    f(DumpableValue::Float(pyfloat.to_f64()))
54                }
55                ref pycomplex @ PyComplex => {
56                    f(DumpableValue::Complex(pycomplex.as_complex()))
57                }
58                ref pystr @ PyStr => {
59                    f(DumpableValue::Str(pystr.as_wtf8()))
60                }
61                ref pylist @ PyList => {
62                    f(DumpableValue::List(&pylist.borrow_vec()))
63                }
64                ref pyset @ PySet => {
65                    let elements = pyset.elements();
66                    f(DumpableValue::Set(&elements))
67                }
68                ref pyfrozen @ PyFrozenSet => {
69                    let elements = pyfrozen.elements();
70                    f(DumpableValue::Frozenset(&elements))
71                }
72                ref pytuple @ PyTuple => {
73                    f(DumpableValue::Tuple(pytuple.as_slice()))
74                }
75                ref pydict @ PyDict => {
76                    let entries = pydict.into_iter().collect::<Vec<_>>();
77                    f(DumpableValue::Dict(&entries))
78                }
79                ref bytes @ PyBytes => {
80                    f(DumpableValue::Bytes(bytes.as_bytes()))
81                }
82                ref bytes @ PyByteArray => {
83                    f(DumpableValue::Bytes(&bytes.borrow_buf()))
84                }
85                ref co @ PyCode => {
86                    f(DumpableValue::Code(co))
87                }
88                _ => return Err(DumpError),
89            });
90            Ok(ret)
91        }
92    }
93
94    #[derive(FromArgs)]
95    struct DumpsArgs {
96        #[pyarg(positional)]
97        value: PyObjectRef,
98        #[pyarg(positional, default = 5)]
99        version: i32,
100        #[pyarg(named, default = true)]
101        allow_code: bool,
102    }
103
104    #[pyfunction]
105    fn dumps(args: DumpsArgs, vm: &VirtualMachine) -> PyResult<PyBytes> {
106        let DumpsArgs {
107            value,
108            allow_code,
109            version,
110        } = args;
111
112        vm.audit("marshal.dumps", || (value.clone(), version))?;
113
114        check_exact_type(&value, vm)?;
115        let mut buf = Vec::new();
116        let mut refs = if version >= 3 {
117            Some(WriterRefTable::new())
118        } else {
119            None
120        };
121        write_object(&mut buf, &value, &mut refs, version, allow_code, vm)?;
122        Ok(PyBytes::from(buf))
123    }
124
125    struct WriterRefEntry {
126        idx: u32,
127        /// Set between `reserve` and `complete` for the object kinds whose
128        /// immutable representation cannot be rebuilt from a back-reference.
129        incomplete: bool,
130    }
131
132    struct WriterRefTable {
133        map: std::collections::HashMap<usize, WriterRefEntry>,
134        next_idx: u32,
135    }
136
137    impl WriterRefTable {
138        fn new() -> Self {
139            Self {
140                map: std::collections::HashMap::new(),
141                next_idx: 0,
142            }
143        }
144        /// `w_ref`: write a back-reference to an object already in the table.
145        /// Reaching an entry that is still being written is a recursion the
146        /// reader could not rebuild, so it is an error rather than a `TYPE_REF`.
147        fn try_ref(&mut self, buf: &mut Vec<u8>, obj: &PyObject) -> Result<bool, ()> {
148            use marshal::Write;
149            let Some(entry) = self.map.get(&obj.get_id()) else {
150                return Ok(false);
151            };
152            if entry.incomplete {
153                return Err(());
154            }
155            buf.write_u8(b'r');
156            buf.write_u32(entry.idx);
157            Ok(true)
158        }
159        fn reserve(&mut self, obj: &PyObject, incomplete: bool) -> u32 {
160            let idx = self.next_idx;
161            self.map
162                .insert(obj.get_id(), WriterRefEntry { idx, incomplete });
163            self.next_idx += 1;
164            idx
165        }
166        /// `w_complete`: the object's contents are on the stream, so a later
167        /// occurrence may reference it.
168        fn complete(&mut self, obj: &PyObject) {
169            if let Some(entry) = self.map.get_mut(&obj.get_id()) {
170                entry.incomplete = false;
171            }
172        }
173    }
174
175    fn write_object(
176        buf: &mut Vec<u8>,
177        obj: &PyObject,
178        refs: &mut Option<WriterRefTable>,
179        version: i32,
180        allow_code: bool,
181        vm: &VirtualMachine,
182    ) -> PyResult<()> {
183        write_object_depth(
184            buf,
185            obj,
186            refs,
187            version,
188            allow_code,
189            vm,
190            marshal::MAX_MARSHAL_STACK_DEPTH,
191        )
192    }
193
194    /// Write a float the way the pre-binary formats carry one: seventeen
195    /// significant digits, prefixed with their length.
196    fn write_float_str(buf: &mut Vec<u8>, value: f64) {
197        use marshal::Write;
198        let digits = rustpython_literal::float::format_general(
199            17,
200            value.abs(),
201            rustpython_literal::format::Case::Lower,
202            false,
203            false,
204        );
205        // A nan carries no sign of its own.
206        let sign = if value.is_sign_negative() && !value.is_nan() {
207            "-"
208        } else {
209            ""
210        };
211        buf.write_u8((sign.len() + digits.len()) as u8);
212        buf.write_slice(sign.as_bytes());
213        buf.write_slice(digits.as_bytes());
214    }
215
216    fn write_object_depth(
217        buf: &mut Vec<u8>,
218        obj: &PyObject,
219        refs: &mut Option<WriterRefTable>,
220        version: i32,
221        allow_code: bool,
222        vm: &VirtualMachine,
223        depth: usize,
224    ) -> PyResult<()> {
225        use marshal::Write;
226        if depth == 0 {
227            return Err(vm.new_value_error("object too deeply nested to marshal"));
228        }
229
230        // Singletons: no FLAG_REF needed
231        let is_singleton = vm.is_none(obj)
232            || obj.class().is(PyBool::static_type())
233            || obj.is(PyStopIteration::static_type())
234            || obj.downcast_ref::<crate::builtins::PyEllipsis>().is_some();
235
236        // FLAG_REF: check if already written, otherwise reserve slot
237        if !is_singleton && let Some(rt) = refs.as_mut() {
238            match rt.try_ref(buf, obj) {
239                Ok(true) => return Ok(()),
240                Ok(false) => {}
241                Err(()) => {
242                    return Err(vm.new_value_error(format!(
243                        "cannot marshal recursion {} objects",
244                        obj.class().name()
245                    )));
246                }
247            }
248        }
249        let type_pos = buf.len();
250        let use_ref = refs.is_some() && !is_singleton;
251        // A code or slice entry stays incomplete until its contents are
252        // written: the reader rebuilds both from their fields, so a
253        // back-reference issued while those fields are still being emitted
254        // would name an object that does not exist yet.
255        let requires_completion = obj.downcast_ref::<PyCode>().is_some()
256            || obj.downcast_ref::<crate::builtins::PySlice>().is_some();
257        if use_ref {
258            refs.as_mut().unwrap().reserve(obj, requires_completion);
259        }
260
261        if vm.is_none(obj) {
262            buf.write_u8(b'N');
263        } else if obj.is(PyStopIteration::static_type()) {
264            buf.write_u8(b'S');
265        } else if obj.class().is(PyBool::static_type()) {
266            let val = obj
267                .downcast_ref::<PyInt>()
268                .is_some_and(|i| !i.as_bigint().is_zero());
269            buf.write_u8(if val { b'T' } else { b'F' });
270        } else if obj.downcast_ref::<crate::builtins::PyEllipsis>().is_some() {
271            buf.write_u8(b'.');
272        } else if let Some(i) = obj.downcast_ref::<PyInt>() {
273            // TYPE_INT for i32 range, TYPE_LONG for larger
274            if let Ok(val) = i32::try_from(i.as_bigint()) {
275                buf.write_u8(b'i');
276                buf.write_u32(val as u32);
277            } else {
278                buf.write_u8(b'l');
279                let (sign, raw) = i.as_bigint().to_bytes_le();
280                let mut digits = Vec::new();
281                let mut accum: u32 = 0;
282                let mut bits = 0u32;
283                for &byte in &raw {
284                    accum |= (byte as u32) << bits;
285                    bits += 8;
286                    while bits >= 15 {
287                        digits.push((accum & 0x7fff) as u16);
288                        accum >>= 15;
289                        bits -= 15;
290                    }
291                }
292                if accum > 0 || digits.is_empty() {
293                    digits.push(accum as u16);
294                }
295                while digits.len() > 1 && *digits.last().unwrap() == 0 {
296                    digits.pop();
297                }
298                let n = digits.len() as i32;
299                let n = if sign == malachite_bigint::Sign::Minus {
300                    -n
301                } else {
302                    n
303                };
304                buf.write_u32(n as u32);
305                for d in &digits {
306                    buf.write_u16(*d);
307                }
308            }
309        } else if let Some(f) = obj.downcast_ref::<PyFloat>() {
310            if version > 1 {
311                buf.write_u8(b'g');
312                buf.write_u64(f.to_f64().to_bits());
313            } else {
314                buf.write_u8(b'f');
315                write_float_str(buf, f.to_f64());
316            }
317        } else if let Some(c) = obj.downcast_ref::<PyComplex>() {
318            let cv = c.as_complex();
319            if version > 1 {
320                buf.write_u8(b'y');
321                buf.write_u64(cv.re.to_bits());
322                buf.write_u64(cv.im.to_bits());
323            } else {
324                buf.write_u8(b'x');
325                write_float_str(buf, cv.re);
326                write_float_str(buf, cv.im);
327            }
328        } else if let Some(s) = obj.downcast_ref::<PyStr>() {
329            let bytes = s.as_wtf8().as_bytes();
330            // Only the formats that carry back-references tell an interned
331            // string apart, and only those from 4 on have the ascii-only forms.
332            let interned = version >= 3 && obj.is_interned();
333            if version >= 4 && bytes.is_ascii() {
334                if bytes.len() <= 255 {
335                    buf.write_u8(if interned { b'Z' } else { b'z' });
336                    buf.write_u8(bytes.len() as u8);
337                } else {
338                    buf.write_u8(if interned { b'A' } else { b'a' });
339                    buf.write_u32(bytes.len() as u32);
340                }
341            } else {
342                buf.write_u8(if interned { b't' } else { b'u' });
343                buf.write_u32(bytes.len() as u32);
344            }
345            buf.write_slice(bytes);
346        } else if let Some(b) = obj.downcast_ref::<PyBytes>() {
347            buf.write_u8(b's');
348            let data = b.as_bytes();
349            buf.write_u32(data.len() as u32);
350            buf.write_slice(data);
351        } else if let Some(b) = obj.downcast_ref::<PyByteArray>() {
352            buf.write_u8(b's');
353            let data = b.borrow_buf();
354            buf.write_u32(data.len() as u32);
355            buf.write_slice(&data);
356        } else if let Some(t) = obj.downcast_ref::<PyTuple>() {
357            // From 4 on a short tuple carries its length in a single byte.
358            if version >= 4 && t.as_slice().len() < 256 {
359                buf.write_u8(b')');
360                buf.write_u8(t.as_slice().len() as u8);
361            } else {
362                buf.write_u8(b'(');
363                buf.write_u32(t.as_slice().len() as u32);
364            }
365            for elem in t.as_slice() {
366                write_object_depth(buf, elem, refs, version, allow_code, vm, depth - 1)?;
367            }
368        } else if let Some(l) = obj.downcast_ref::<PyList>() {
369            buf.write_u8(b'[');
370            let items = l.borrow_vec();
371            buf.write_u32(items.len() as u32);
372            for elem in items.iter() {
373                write_object_depth(buf, elem, refs, version, allow_code, vm, depth - 1)?;
374            }
375        } else if let Some(d) = obj.downcast_ref::<PyDict>() {
376            buf.write_u8(b'{');
377            for (k, v) in d {
378                write_object_depth(buf, &k, refs, version, allow_code, vm, depth - 1)?;
379                write_object_depth(buf, &v, refs, version, allow_code, vm, depth - 1)?;
380            }
381            buf.write_u8(b'0'); // TYPE_NULL terminator
382        } else if let Some(s) = obj.downcast_ref::<PySet>() {
383            buf.write_u8(b'<');
384            write_set_elements(buf, &s.elements(), refs, version, allow_code, vm, depth)?;
385        } else if let Some(s) = obj.downcast_ref::<PyFrozenSet>() {
386            buf.write_u8(b'>');
387            write_set_elements(buf, &s.elements(), refs, version, allow_code, vm, depth)?;
388        } else if let Some(co) = obj.downcast_ref::<PyCode>() {
389            if !allow_code {
390                return Err(vm.new_value_error("marshalling code objects is disallowed"));
391            }
392            buf.write_u8(b'c');
393            // `Literal` holds the exact object a constant was built from, so
394            // route `co_consts` back through the object writer: it reaches the
395            // values `BorrowedConstant` cannot describe and shares the one
396            // reference table the reader indexes against.
397            marshal::serialize_code_with(buf, &co.code, |buf, constant| {
398                let constant = PyObjectRef::from(constant.clone());
399                write_object_depth(buf, &constant, refs, version, allow_code, vm, depth - 1)
400            })?;
401        } else if let Some(sl) = obj.downcast_ref::<crate::builtins::PySlice>() {
402            if version < 5 {
403                return Err(vm.new_value_error("unmarshallable object"));
404            }
405            buf.write_u8(b':');
406            let none: PyObjectRef = vm.ctx.none();
407            write_object_depth(
408                buf,
409                sl.start.as_ref().unwrap_or(&none),
410                refs,
411                version,
412                allow_code,
413                vm,
414                depth - 1,
415            )?;
416            write_object_depth(buf, &sl.stop, refs, version, allow_code, vm, depth - 1)?;
417            write_object_depth(
418                buf,
419                sl.step.as_ref().unwrap_or(&none),
420                refs,
421                version,
422                allow_code,
423                vm,
424                depth - 1,
425            )?;
426        } else if let Ok(bytes_like) = ArgBytesLike::try_from_object(vm, obj.to_owned()) {
427            buf.write_u8(b's');
428            let data = bytes_like.borrow_buf();
429            buf.write_u32(data.len() as u32);
430            buf.write_slice(&data);
431        } else {
432            return Err(vm.new_value_error("unmarshallable object"));
433        }
434
435        if use_ref {
436            buf[type_pos] |= marshal::FLAG_REF;
437            if requires_completion {
438                refs.as_mut().unwrap().complete(obj);
439            }
440        }
441        Ok(())
442    }
443
444    /// Serialize set elements in `sorted(v, key=marshal.dumps)` order.
445    fn write_set_elements(
446        buf: &mut Vec<u8>,
447        elems: &[PyObjectRef],
448        refs: &mut Option<WriterRefTable>,
449        version: i32,
450        allow_code: bool,
451        vm: &VirtualMachine,
452        depth: usize,
453    ) -> PyResult<()> {
454        use marshal::Write;
455        buf.write_u32(elems.len() as u32);
456        let mut pairs = Vec::with_capacity(elems.len());
457        for elem in elems {
458            let mut dumped = Vec::new();
459            let mut inner_refs = (version >= 3).then(WriterRefTable::new);
460            write_object(&mut dumped, elem, &mut inner_refs, version, allow_code, vm)?;
461            pairs.push((dumped, elem.clone()));
462        }
463        pairs.sort_by(|a, b| a.0.cmp(&b.0));
464        for (_, elem) in &pairs {
465            write_object_depth(buf, elem, refs, version, allow_code, vm, depth - 1)?;
466        }
467        Ok(())
468    }
469
470    #[derive(FromArgs)]
471    struct DumpArgs {
472        #[pyarg(positional)]
473        value: PyObjectRef,
474        #[pyarg(positional)]
475        file: PyObjectRef,
476        #[pyarg(positional, default = 5)]
477        version: i32,
478        #[pyarg(named, default = true)]
479        allow_code: bool,
480    }
481
482    #[pyfunction]
483    fn dump(args: DumpArgs, vm: &VirtualMachine) -> PyResult<()> {
484        let dumped = dumps(
485            DumpsArgs {
486                value: args.value,
487                version: args.version,
488                allow_code: args.allow_code,
489            },
490            vm,
491        )?;
492        vm.call_method(&args.file, "write", (dumped,))?;
493        Ok(())
494    }
495
496    #[derive(Copy, Clone)]
497    struct PyMarshalBag<'a> {
498        vm: &'a VirtualMachine,
499        pending_error: &'a RefCell<Option<PyBaseExceptionRef>>,
500        allow_code: bool,
501    }
502
503    impl<'a> PyMarshalBag<'a> {
504        fn new(
505            vm: &'a VirtualMachine,
506            pending_error: &'a RefCell<Option<PyBaseExceptionRef>>,
507            allow_code: bool,
508        ) -> Self {
509            Self {
510                vm,
511                pending_error,
512                allow_code,
513            }
514        }
515
516        /// Room for a container the decoder publishes before it reads what
517        /// goes in it. The length is the input's to choose, so the room is
518        /// asked for rather than assumed: a length no allocator can serve is
519        /// a MemoryError, not an aborted process.
520        fn placeholder_elements(
521            &self,
522            len: usize,
523        ) -> Result<Vec<PyObjectRef>, marshal::MarshalError> {
524            let mut elements = Vec::new();
525            elements
526                .try_reserve_exact(len)
527                .map_err(|_| self.remember_python_error(self.vm.no_memory_error()))?;
528            elements.resize(len, self.vm.ctx.none());
529            Ok(elements)
530        }
531
532        fn remember_python_error(&self, error: PyBaseExceptionRef) -> marshal::MarshalError {
533            let mut pending = self.pending_error.borrow_mut();
534            if pending.is_none() {
535                *pending = Some(error);
536            }
537            marshal::MarshalError::BadType
538        }
539    }
540
541    impl<'a> marshal::MarshalBag for PyMarshalBag<'a> {
542        type Value = PyObjectRef;
543        type ConstantBag = PyVmBag<'a>;
544
545        fn make_bool(&self, value: bool) -> Self::Value {
546            self.vm.ctx.new_bool(value).into()
547        }
548        fn make_none(&self) -> Self::Value {
549            self.vm.ctx.none()
550        }
551        fn make_ellipsis(&self) -> Self::Value {
552            self.vm.ctx.ellipsis.clone().into()
553        }
554        fn make_float(&self, value: f64) -> Self::Value {
555            self.vm.ctx.new_float(value).into()
556        }
557        fn make_complex(&self, value: num_complex::Complex64) -> Self::Value {
558            self.vm.ctx.new_complex(value).into()
559        }
560        fn make_str(&self, value: &Wtf8) -> Self::Value {
561            self.vm.ctx.new_str(value).into()
562        }
563        fn make_interned_str(&self, value: &Wtf8) -> Self::Value {
564            self.vm.ctx.intern_str(value).to_owned().into()
565        }
566        fn make_bytes(&self, value: &[u8]) -> Self::Value {
567            self.vm.ctx.new_bytes(value.to_vec()).into()
568        }
569        fn make_int(&self, value: BigInt) -> Self::Value {
570            self.vm.ctx.new_int(value).into()
571        }
572        fn make_tuple(&self, elements: impl Iterator<Item = Self::Value>) -> Self::Value {
573            self.vm.ctx.new_tuple(elements.collect()).into()
574        }
575        fn make_tuple_placeholder(
576            &self,
577            len: usize,
578        ) -> Result<Option<Self::Value>, marshal::MarshalError> {
579            let elements = self.placeholder_elements(len)?;
580            Ok(Some(PyTuple::new_ref(elements, &self.vm.ctx).into()))
581        }
582        fn set_tuple_item(
583            &self,
584            tuple: &Self::Value,
585            index: usize,
586            value: Self::Value,
587        ) -> Result<(), marshal::MarshalError> {
588            let tuple = tuple
589                .downcast_ref::<PyTuple>()
590                .ok_or(marshal::MarshalError::BadType)?;
591            // SAFETY: compiler-core calls this only on a fresh placeholder,
592            // once per index, before returning it to Python code.
593            unsafe { tuple.payload.set_marshal_item(index, value) };
594            Ok(())
595        }
596        fn make_code(&self, code: CodeObject) -> Result<Self::Value, marshal::MarshalError> {
597            if !self.allow_code {
598                return Err(self.remember_python_error(
599                    self.vm
600                        .new_value_error("unmarshalling code objects is disallowed"),
601                ));
602            }
603            Ok(crate::builtins::PyCode::new_ref_with_bag(self.vm, code).into())
604        }
605        fn make_stop_iter(&self) -> Result<Self::Value, marshal::MarshalError> {
606            Ok(self.vm.ctx.exceptions.stop_iteration.to_owned().into())
607        }
608        fn make_list(
609            &self,
610            it: impl Iterator<Item = Self::Value>,
611        ) -> Result<Self::Value, marshal::MarshalError> {
612            Ok(self.vm.ctx.new_list(it.collect()).into())
613        }
614        fn make_list_placeholder(
615            &self,
616            len: usize,
617        ) -> Result<Option<Self::Value>, marshal::MarshalError> {
618            let elements = self.placeholder_elements(len)?;
619            Ok(Some(self.vm.ctx.new_list(elements).into()))
620        }
621        fn set_list_item(
622            &self,
623            list: &Self::Value,
624            index: usize,
625            value: Self::Value,
626        ) -> Result<(), marshal::MarshalError> {
627            let list = list
628                .downcast_ref::<PyList>()
629                .ok_or(marshal::MarshalError::BadType)?;
630            list.borrow_vec_mut()[index] = value;
631            Ok(())
632        }
633        fn make_set(
634            &self,
635            it: impl Iterator<Item = Self::Value>,
636        ) -> Result<Self::Value, marshal::MarshalError> {
637            let set = PySet::default().into_ref(&self.vm.ctx);
638            for elem in it {
639                set.add(elem, self.vm)
640                    .map_err(|error| self.remember_python_error(error))?;
641            }
642            Ok(set.into())
643        }
644        fn make_set_placeholder(&self) -> Option<Self::Value> {
645            Some(PySet::default().into_ref(&self.vm.ctx).into())
646        }
647        fn insert_set_item(
648            &self,
649            set: &Self::Value,
650            value: Self::Value,
651        ) -> Result<(), marshal::MarshalError> {
652            let set = set
653                .downcast_ref::<PySet>()
654                .ok_or(marshal::MarshalError::BadType)?;
655            set.add(value, self.vm)
656                .map_err(|error| self.remember_python_error(error))
657        }
658        fn make_frozenset(
659            &self,
660            it: impl Iterator<Item = Self::Value>,
661        ) -> Result<Self::Value, marshal::MarshalError> {
662            PyFrozenSet::from_iter(self.vm, it)
663                .map(|set| set.to_pyobject(self.vm))
664                .map_err(|error| self.remember_python_error(error))
665        }
666        fn make_dict(
667            &self,
668            it: impl Iterator<Item = (Self::Value, Self::Value)>,
669        ) -> Result<Self::Value, marshal::MarshalError> {
670            let dict = self.vm.ctx.new_dict();
671            for (k, v) in it {
672                dict.set_item(&*k, v, self.vm)
673                    .map_err(|error| self.remember_python_error(error))?;
674            }
675            Ok(dict.into())
676        }
677        fn make_dict_placeholder(&self) -> Option<Self::Value> {
678            Some(self.vm.ctx.new_dict().into())
679        }
680        fn insert_dict_item(
681            &self,
682            dict: &Self::Value,
683            key: Self::Value,
684            value: Self::Value,
685        ) -> Result<(), marshal::MarshalError> {
686            let dict = dict
687                .downcast_ref::<PyDict>()
688                .ok_or(marshal::MarshalError::BadType)?;
689            dict.set_item(&*key, value, self.vm)
690                .map_err(|error| self.remember_python_error(error))
691        }
692        fn make_slice(
693            &self,
694            start: Self::Value,
695            stop: Self::Value,
696            step: Self::Value,
697        ) -> Result<Self::Value, marshal::MarshalError> {
698            use crate::builtins::PySlice;
699            let vm = self.vm;
700            Ok(PySlice {
701                start: if vm.is_none(&start) {
702                    None
703                } else {
704                    Some(start)
705                },
706                stop,
707                step: if vm.is_none(&step) { None } else { Some(step) },
708            }
709            .into_ref(&vm.ctx)
710            .into())
711        }
712        fn constant_bag(self) -> Self::ConstantBag {
713            PyVmBag(self.vm)
714        }
715        /// `Literal` wraps any object, so a decoded `co_consts` entry is
716        /// already its own compiler-side constant — no placeholder is needed
717        /// and `make_code_with_constants` keeps the default.
718        fn constant_ref_from_value(&self, value: &Self::Value) -> Option<Literal> {
719            Some(Literal::from(value.clone()))
720        }
721        fn bytes_from_value(&self, value: &Self::Value) -> Option<Vec<u8>> {
722            value
723                .downcast_ref::<PyBytes>()
724                .map(|bytes| bytes.as_bytes().to_vec())
725        }
726        fn str_from_value(&self, value: &Self::Value) -> Option<String> {
727            value
728                .downcast_ref::<PyStr>()
729                .map(|str| str.to_string_lossy().into_owned())
730        }
731        fn tuple_elements_from_value(&self, value: &Self::Value) -> Option<Vec<Self::Value>> {
732            value
733                .downcast_ref::<PyTuple>()
734                .map(|tuple| tuple.as_slice().to_vec())
735        }
736    }
737
738    fn deserialize_value(
739        rdr: &mut impl marshal::Read,
740        allow_code: bool,
741        vm: &VirtualMachine,
742    ) -> PyResult<PyObjectRef> {
743        let pending_error = RefCell::new(None);
744        match marshal::deserialize_value(rdr, PyMarshalBag::new(vm, &pending_error, allow_code)) {
745            Ok(value) => Ok(value),
746            Err(error) => Err(pending_error.into_inner().unwrap_or_else(|| match error {
747                marshal::MarshalError::Eof => vm.new_eof_error("EOF read where not expected"),
748                marshal::MarshalError::EofObject => {
749                    vm.new_eof_error("EOF read where object expected")
750                }
751                marshal::MarshalError::DataTooShort => vm.new_eof_error("marshal data too short"),
752                error @ marshal::MarshalError::NullObject => vm.new_type_error(error.to_string()),
753                error @ (marshal::MarshalError::BadSize(_)
754                | marshal::MarshalError::UnknownType
755                | marshal::MarshalError::InvalidRef) => {
756                    vm.new_value_error(format!("bad marshal data ({error})"))
757                }
758                _ => vm.new_value_error("bad marshal data"),
759            })),
760        }
761    }
762
763    #[derive(FromArgs)]
764    struct LoadsArgs {
765        #[pyarg(positional)]
766        // marshal_loads_impl takes `bytes: Py_buffer`, a y* argument.
767        bytes: ArgBytesLike,
768        #[pyarg(named, default = true)]
769        allow_code: bool,
770    }
771
772    #[pyfunction]
773    fn loads(args: LoadsArgs, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
774        let LoadsArgs { bytes, allow_code } = args;
775        let buf = bytes.borrow_buf();
776
777        deserialize_value(&mut &buf[..], allow_code, vm)
778    }
779
780    #[derive(FromArgs)]
781    struct LoadArgs {
782        #[pyarg(positional)]
783        file: PyObjectRef,
784        #[pyarg(named, default = true)]
785        allow_code: bool,
786    }
787
788    #[pyfunction]
789    fn load(args: LoadArgs, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
790        let mut rdr = ReadableFile {
791            file: args.file,
792            vm,
793            buf: Vec::new(),
794            error: None,
795        };
796        let pending_error = RefCell::new(None);
797        let result = marshal::deserialize_value(
798            &mut rdr,
799            PyMarshalBag::new(vm, &pending_error, args.allow_code),
800        );
801        if let Some(err) = rdr.error.take() {
802            return Err(err);
803        }
804        match result {
805            Ok(value) => Ok(value),
806            Err(error) => Err(pending_error.into_inner().unwrap_or_else(|| match error {
807                marshal::MarshalError::Eof => vm.new_eof_error("EOF read where not expected"),
808                marshal::MarshalError::EofObject => {
809                    vm.new_eof_error("EOF read where object expected")
810                }
811                marshal::MarshalError::DataTooShort => vm.new_eof_error("marshal data too short"),
812                error @ marshal::MarshalError::NullObject => vm.new_type_error(error.to_string()),
813                error @ (marshal::MarshalError::BadSize(_)
814                | marshal::MarshalError::UnknownType
815                | marshal::MarshalError::InvalidRef) => {
816                    vm.new_value_error(format!("bad marshal data ({error})"))
817                }
818                _ => vm.new_value_error("bad marshal data"),
819            })),
820        }
821    }
822
823    /// File-backed marshal reader. `r_string` fills a scratch buffer via
824    /// `readinto` on a contiguous writable memoryview.
825    struct ReadableFile<'a> {
826        file: PyObjectRef,
827        vm: &'a VirtualMachine,
828        buf: Vec<u8>,
829        error: Option<PyBaseExceptionRef>,
830    }
831
832    impl ReadableFile<'_> {
833        fn r_string(&mut self, n: usize) -> PyResult<()> {
834            self.buf.clear();
835            let bytearray = PyByteArray::from(vec![0u8; n]).into_ref(&self.vm.ctx);
836            let memoryview = PyMemoryView::from_object_with_flags(
837                bytearray.as_object(),
838                crate::protocol::BufferFlags::CONTIG,
839                self.vm,
840            )?
841            .into_ref(&self.vm.ctx);
842            let nread_obj = self.vm.call_method(&self.file, "readinto", (memoryview,))?;
843            let nread = nread_obj
844                .try_index(self.vm)?
845                .try_to_primitive::<isize>(self.vm)?;
846            let n_isize = isize::try_from(n).unwrap_or(isize::MAX);
847            if nread != n_isize {
848                if nread > n_isize {
849                    return Err(self.vm.new_value_error(format!(
850                        "read() returned too much data: {n} bytes requested, {nread} returned"
851                    )));
852                }
853                return Err(self.vm.new_eof_error("EOF read where not expected"));
854            }
855            self.buf.extend_from_slice(&bytearray.borrow_buf());
856            Ok(())
857        }
858    }
859
860    impl marshal::Read for ReadableFile<'_> {
861        fn read_slice(&mut self, n: u32) -> Result<&[u8], marshal::MarshalError> {
862            if self.error.is_some() {
863                return Err(marshal::MarshalError::Eof);
864            }
865            if let Err(e) = self.r_string(n as usize) {
866                self.error = Some(e);
867                return Err(marshal::MarshalError::Eof);
868            }
869            Ok(&self.buf)
870        }
871    }
872
873    /// Reject subclasses of marshallable types (int, float, complex, tuple, etc.).
874    fn check_exact_type(obj: &PyObject, vm: &VirtualMachine) -> PyResult<()> {
875        let cls = obj.class();
876        // bool is a subclass of int but is marshallable
877        if cls.is(PyBool::static_type()) {
878            return Ok(());
879        }
880        for base in [
881            PyInt::static_type(),
882            PyFloat::static_type(),
883            PyComplex::static_type(),
884            PyTuple::static_type(),
885            PyList::static_type(),
886            PyDict::static_type(),
887            PySet::static_type(),
888            PyFrozenSet::static_type(),
889        ] {
890            if cls.fast_issubclass(base) && !cls.is(base) {
891                return Err(vm.new_value_error("unmarshallable object"));
892            }
893        }
894        Ok(())
895    }
896}