1pub(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 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 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 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 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 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 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 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 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 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 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 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'); } 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 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 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 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 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 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 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 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 fn check_exact_type(obj: &PyObject, vm: &VirtualMachine) -> PyResult<()> {
875 let cls = obj.class();
876 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}