pub(crate) use decl::module_def;
#[pymodule(name = "marshal")]
mod decl {
use crate::builtins::code::{CodeObject, Literal, PyVmBag};
use crate::class::StaticType;
use crate::common::wtf8::Wtf8;
use crate::{
PyObject, PyObjectRef, PyResult, TryFromObject, VirtualMachine,
builtins::{
PyBaseExceptionRef, PyBool, PyByteArray, PyBytes, PyCode, PyComplex, PyDict,
PyEllipsis, PyFloat, PyFrozenSet, PyInt, PyList, PyMemoryView, PyNone, PySet,
PyStopIteration, PyStr, PyTuple,
},
convert::ToPyObject,
function::ArgBytesLike,
object::{AsObject, PyPayload},
};
use core::cell::RefCell;
use malachite_bigint::BigInt;
use num_traits::Zero;
use rustpython_compiler_core::marshal::{self, DumpableValue};
#[pyattr(name = "version")]
use marshal::FORMAT_VERSION;
pub struct DumpError;
impl marshal::Dumpable for PyObjectRef {
type Error = DumpError;
type Constant = Literal;
fn with_dump<R>(
&self,
f: impl FnOnce(DumpableValue<'_, Self>) -> R,
) -> Result<R, Self::Error> {
if self.is(PyStopIteration::static_type()) {
return Ok(f(DumpableValue::StopIter));
}
let ret = match_class!(match self {
PyNone => f(DumpableValue::None),
PyEllipsis => f(DumpableValue::Ellipsis),
ref pyint @ PyInt => {
if self.class().is(PyBool::static_type()) {
f(DumpableValue::Boolean(!pyint.as_bigint().is_zero()))
} else {
f(DumpableValue::Integer(pyint.as_bigint()))
}
}
ref pyfloat @ PyFloat => {
f(DumpableValue::Float(pyfloat.to_f64()))
}
ref pycomplex @ PyComplex => {
f(DumpableValue::Complex(pycomplex.as_complex()))
}
ref pystr @ PyStr => {
f(DumpableValue::Str(pystr.as_wtf8()))
}
ref pylist @ PyList => {
f(DumpableValue::List(&pylist.borrow_vec()))
}
ref pyset @ PySet => {
let elements = pyset.elements();
f(DumpableValue::Set(&elements))
}
ref pyfrozen @ PyFrozenSet => {
let elements = pyfrozen.elements();
f(DumpableValue::Frozenset(&elements))
}
ref pytuple @ PyTuple => {
f(DumpableValue::Tuple(pytuple.as_slice()))
}
ref pydict @ PyDict => {
let entries = pydict.into_iter().collect::<Vec<_>>();
f(DumpableValue::Dict(&entries))
}
ref bytes @ PyBytes => {
f(DumpableValue::Bytes(bytes.as_bytes()))
}
ref bytes @ PyByteArray => {
f(DumpableValue::Bytes(&bytes.borrow_buf()))
}
ref co @ PyCode => {
f(DumpableValue::Code(co))
}
_ => return Err(DumpError),
});
Ok(ret)
}
}
#[derive(FromArgs)]
struct DumpsArgs {
#[pyarg(positional)]
value: PyObjectRef,
#[pyarg(positional, default = 5)]
version: i32,
#[pyarg(named, default = true)]
allow_code: bool,
}
#[pyfunction]
fn dumps(args: DumpsArgs, vm: &VirtualMachine) -> PyResult<PyBytes> {
let DumpsArgs {
value,
allow_code,
version,
} = args;
vm.audit("marshal.dumps", || (value.clone(), version))?;
check_exact_type(&value, vm)?;
let mut buf = Vec::new();
let mut refs = if version >= 3 {
Some(WriterRefTable::new())
} else {
None
};
write_object(&mut buf, &value, &mut refs, version, allow_code, vm)?;
Ok(PyBytes::from(buf))
}
struct WriterRefEntry {
idx: u32,
incomplete: bool,
}
struct WriterRefTable {
map: std::collections::HashMap<usize, WriterRefEntry>,
next_idx: u32,
}
impl WriterRefTable {
fn new() -> Self {
Self {
map: std::collections::HashMap::new(),
next_idx: 0,
}
}
fn try_ref(&mut self, buf: &mut Vec<u8>, obj: &PyObject) -> Result<bool, ()> {
use marshal::Write;
let Some(entry) = self.map.get(&obj.get_id()) else {
return Ok(false);
};
if entry.incomplete {
return Err(());
}
buf.write_u8(b'r');
buf.write_u32(entry.idx);
Ok(true)
}
fn reserve(&mut self, obj: &PyObject, incomplete: bool) -> u32 {
let idx = self.next_idx;
self.map
.insert(obj.get_id(), WriterRefEntry { idx, incomplete });
self.next_idx += 1;
idx
}
fn complete(&mut self, obj: &PyObject) {
if let Some(entry) = self.map.get_mut(&obj.get_id()) {
entry.incomplete = false;
}
}
}
fn write_object(
buf: &mut Vec<u8>,
obj: &PyObject,
refs: &mut Option<WriterRefTable>,
version: i32,
allow_code: bool,
vm: &VirtualMachine,
) -> PyResult<()> {
write_object_depth(
buf,
obj,
refs,
version,
allow_code,
vm,
marshal::MAX_MARSHAL_STACK_DEPTH,
)
}
fn write_float_str(buf: &mut Vec<u8>, value: f64) {
use marshal::Write;
let digits = rustpython_literal::float::format_general(
17,
value.abs(),
rustpython_literal::format::Case::Lower,
false,
false,
);
let sign = if value.is_sign_negative() && !value.is_nan() {
"-"
} else {
""
};
buf.write_u8((sign.len() + digits.len()) as u8);
buf.write_slice(sign.as_bytes());
buf.write_slice(digits.as_bytes());
}
fn write_object_depth(
buf: &mut Vec<u8>,
obj: &PyObject,
refs: &mut Option<WriterRefTable>,
version: i32,
allow_code: bool,
vm: &VirtualMachine,
depth: usize,
) -> PyResult<()> {
use marshal::Write;
if depth == 0 {
return Err(vm.new_value_error("object too deeply nested to marshal"));
}
let is_singleton = vm.is_none(obj)
|| obj.class().is(PyBool::static_type())
|| obj.is(PyStopIteration::static_type())
|| obj.downcast_ref::<crate::builtins::PyEllipsis>().is_some();
if !is_singleton && let Some(rt) = refs.as_mut() {
match rt.try_ref(buf, obj) {
Ok(true) => return Ok(()),
Ok(false) => {}
Err(()) => {
return Err(vm.new_value_error(format!(
"cannot marshal recursion {} objects",
obj.class().name()
)));
}
}
}
let type_pos = buf.len();
let use_ref = refs.is_some() && !is_singleton;
let requires_completion = obj.downcast_ref::<PyCode>().is_some()
|| obj.downcast_ref::<crate::builtins::PySlice>().is_some();
if use_ref {
refs.as_mut().unwrap().reserve(obj, requires_completion);
}
if vm.is_none(obj) {
buf.write_u8(b'N');
} else if obj.is(PyStopIteration::static_type()) {
buf.write_u8(b'S');
} else if obj.class().is(PyBool::static_type()) {
let val = obj
.downcast_ref::<PyInt>()
.is_some_and(|i| !i.as_bigint().is_zero());
buf.write_u8(if val { b'T' } else { b'F' });
} else if obj.downcast_ref::<crate::builtins::PyEllipsis>().is_some() {
buf.write_u8(b'.');
} else if let Some(i) = obj.downcast_ref::<PyInt>() {
if let Ok(val) = i32::try_from(i.as_bigint()) {
buf.write_u8(b'i');
buf.write_u32(val as u32);
} else {
buf.write_u8(b'l');
let (sign, raw) = i.as_bigint().to_bytes_le();
let mut digits = Vec::new();
let mut accum: u32 = 0;
let mut bits = 0u32;
for &byte in &raw {
accum |= (byte as u32) << bits;
bits += 8;
while bits >= 15 {
digits.push((accum & 0x7fff) as u16);
accum >>= 15;
bits -= 15;
}
}
if accum > 0 || digits.is_empty() {
digits.push(accum as u16);
}
while digits.len() > 1 && *digits.last().unwrap() == 0 {
digits.pop();
}
let n = digits.len() as i32;
let n = if sign == malachite_bigint::Sign::Minus {
-n
} else {
n
};
buf.write_u32(n as u32);
for d in &digits {
buf.write_u16(*d);
}
}
} else if let Some(f) = obj.downcast_ref::<PyFloat>() {
if version > 1 {
buf.write_u8(b'g');
buf.write_u64(f.to_f64().to_bits());
} else {
buf.write_u8(b'f');
write_float_str(buf, f.to_f64());
}
} else if let Some(c) = obj.downcast_ref::<PyComplex>() {
let cv = c.as_complex();
if version > 1 {
buf.write_u8(b'y');
buf.write_u64(cv.re.to_bits());
buf.write_u64(cv.im.to_bits());
} else {
buf.write_u8(b'x');
write_float_str(buf, cv.re);
write_float_str(buf, cv.im);
}
} else if let Some(s) = obj.downcast_ref::<PyStr>() {
let bytes = s.as_wtf8().as_bytes();
let interned = version >= 3 && obj.is_interned();
if version >= 4 && bytes.is_ascii() {
if bytes.len() <= 255 {
buf.write_u8(if interned { b'Z' } else { b'z' });
buf.write_u8(bytes.len() as u8);
} else {
buf.write_u8(if interned { b'A' } else { b'a' });
buf.write_u32(bytes.len() as u32);
}
} else {
buf.write_u8(if interned { b't' } else { b'u' });
buf.write_u32(bytes.len() as u32);
}
buf.write_slice(bytes);
} else if let Some(b) = obj.downcast_ref::<PyBytes>() {
buf.write_u8(b's');
let data = b.as_bytes();
buf.write_u32(data.len() as u32);
buf.write_slice(data);
} else if let Some(b) = obj.downcast_ref::<PyByteArray>() {
buf.write_u8(b's');
let data = b.borrow_buf();
buf.write_u32(data.len() as u32);
buf.write_slice(&data);
} else if let Some(t) = obj.downcast_ref::<PyTuple>() {
if version >= 4 && t.as_slice().len() < 256 {
buf.write_u8(b')');
buf.write_u8(t.as_slice().len() as u8);
} else {
buf.write_u8(b'(');
buf.write_u32(t.as_slice().len() as u32);
}
for elem in t.as_slice() {
write_object_depth(buf, elem, refs, version, allow_code, vm, depth - 1)?;
}
} else if let Some(l) = obj.downcast_ref::<PyList>() {
buf.write_u8(b'[');
let items = l.borrow_vec();
buf.write_u32(items.len() as u32);
for elem in items.iter() {
write_object_depth(buf, elem, refs, version, allow_code, vm, depth - 1)?;
}
} else if let Some(d) = obj.downcast_ref::<PyDict>() {
buf.write_u8(b'{');
for (k, v) in d {
write_object_depth(buf, &k, refs, version, allow_code, vm, depth - 1)?;
write_object_depth(buf, &v, refs, version, allow_code, vm, depth - 1)?;
}
buf.write_u8(b'0'); } else if let Some(s) = obj.downcast_ref::<PySet>() {
buf.write_u8(b'<');
write_set_elements(buf, &s.elements(), refs, version, allow_code, vm, depth)?;
} else if let Some(s) = obj.downcast_ref::<PyFrozenSet>() {
buf.write_u8(b'>');
write_set_elements(buf, &s.elements(), refs, version, allow_code, vm, depth)?;
} else if let Some(co) = obj.downcast_ref::<PyCode>() {
if !allow_code {
return Err(vm.new_value_error("marshalling code objects is disallowed"));
}
buf.write_u8(b'c');
marshal::serialize_code_with(buf, &co.code, |buf, constant| {
let constant = PyObjectRef::from(constant.clone());
write_object_depth(buf, &constant, refs, version, allow_code, vm, depth - 1)
})?;
} else if let Some(sl) = obj.downcast_ref::<crate::builtins::PySlice>() {
if version < 5 {
return Err(vm.new_value_error("unmarshallable object"));
}
buf.write_u8(b':');
let none: PyObjectRef = vm.ctx.none();
write_object_depth(
buf,
sl.start.as_ref().unwrap_or(&none),
refs,
version,
allow_code,
vm,
depth - 1,
)?;
write_object_depth(buf, &sl.stop, refs, version, allow_code, vm, depth - 1)?;
write_object_depth(
buf,
sl.step.as_ref().unwrap_or(&none),
refs,
version,
allow_code,
vm,
depth - 1,
)?;
} else if let Ok(bytes_like) = ArgBytesLike::try_from_object(vm, obj.to_owned()) {
buf.write_u8(b's');
let data = bytes_like.borrow_buf();
buf.write_u32(data.len() as u32);
buf.write_slice(&data);
} else {
return Err(vm.new_value_error("unmarshallable object"));
}
if use_ref {
buf[type_pos] |= marshal::FLAG_REF;
if requires_completion {
refs.as_mut().unwrap().complete(obj);
}
}
Ok(())
}
fn write_set_elements(
buf: &mut Vec<u8>,
elems: &[PyObjectRef],
refs: &mut Option<WriterRefTable>,
version: i32,
allow_code: bool,
vm: &VirtualMachine,
depth: usize,
) -> PyResult<()> {
use marshal::Write;
buf.write_u32(elems.len() as u32);
let mut pairs = Vec::with_capacity(elems.len());
for elem in elems {
let mut dumped = Vec::new();
let mut inner_refs = (version >= 3).then(WriterRefTable::new);
write_object(&mut dumped, elem, &mut inner_refs, version, allow_code, vm)?;
pairs.push((dumped, elem.clone()));
}
pairs.sort_by(|a, b| a.0.cmp(&b.0));
for (_, elem) in &pairs {
write_object_depth(buf, elem, refs, version, allow_code, vm, depth - 1)?;
}
Ok(())
}
#[derive(FromArgs)]
struct DumpArgs {
#[pyarg(positional)]
value: PyObjectRef,
#[pyarg(positional)]
file: PyObjectRef,
#[pyarg(positional, default = 5)]
version: i32,
#[pyarg(named, default = true)]
allow_code: bool,
}
#[pyfunction]
fn dump(args: DumpArgs, vm: &VirtualMachine) -> PyResult<()> {
let dumped = dumps(
DumpsArgs {
value: args.value,
version: args.version,
allow_code: args.allow_code,
},
vm,
)?;
vm.call_method(&args.file, "write", (dumped,))?;
Ok(())
}
#[derive(Copy, Clone)]
struct PyMarshalBag<'a> {
vm: &'a VirtualMachine,
pending_error: &'a RefCell<Option<PyBaseExceptionRef>>,
allow_code: bool,
}
impl<'a> PyMarshalBag<'a> {
fn new(
vm: &'a VirtualMachine,
pending_error: &'a RefCell<Option<PyBaseExceptionRef>>,
allow_code: bool,
) -> Self {
Self {
vm,
pending_error,
allow_code,
}
}
fn placeholder_elements(
&self,
len: usize,
) -> Result<Vec<PyObjectRef>, marshal::MarshalError> {
let mut elements = Vec::new();
elements
.try_reserve_exact(len)
.map_err(|_| self.remember_python_error(self.vm.no_memory_error()))?;
elements.resize(len, self.vm.ctx.none());
Ok(elements)
}
fn remember_python_error(&self, error: PyBaseExceptionRef) -> marshal::MarshalError {
let mut pending = self.pending_error.borrow_mut();
if pending.is_none() {
*pending = Some(error);
}
marshal::MarshalError::BadType
}
}
impl<'a> marshal::MarshalBag for PyMarshalBag<'a> {
type Value = PyObjectRef;
type ConstantBag = PyVmBag<'a>;
fn make_bool(&self, value: bool) -> Self::Value {
self.vm.ctx.new_bool(value).into()
}
fn make_none(&self) -> Self::Value {
self.vm.ctx.none()
}
fn make_ellipsis(&self) -> Self::Value {
self.vm.ctx.ellipsis.clone().into()
}
fn make_float(&self, value: f64) -> Self::Value {
self.vm.ctx.new_float(value).into()
}
fn make_complex(&self, value: num_complex::Complex64) -> Self::Value {
self.vm.ctx.new_complex(value).into()
}
fn make_str(&self, value: &Wtf8) -> Self::Value {
self.vm.ctx.new_str(value).into()
}
fn make_interned_str(&self, value: &Wtf8) -> Self::Value {
self.vm.ctx.intern_str(value).to_owned().into()
}
fn make_bytes(&self, value: &[u8]) -> Self::Value {
self.vm.ctx.new_bytes(value.to_vec()).into()
}
fn make_int(&self, value: BigInt) -> Self::Value {
self.vm.ctx.new_int(value).into()
}
fn make_tuple(&self, elements: impl Iterator<Item = Self::Value>) -> Self::Value {
self.vm.ctx.new_tuple(elements.collect()).into()
}
fn make_tuple_placeholder(
&self,
len: usize,
) -> Result<Option<Self::Value>, marshal::MarshalError> {
let elements = self.placeholder_elements(len)?;
Ok(Some(PyTuple::new_ref(elements, &self.vm.ctx).into()))
}
fn set_tuple_item(
&self,
tuple: &Self::Value,
index: usize,
value: Self::Value,
) -> Result<(), marshal::MarshalError> {
let tuple = tuple
.downcast_ref::<PyTuple>()
.ok_or(marshal::MarshalError::BadType)?;
unsafe { tuple.payload.set_marshal_item(index, value) };
Ok(())
}
fn make_code(&self, code: CodeObject) -> Result<Self::Value, marshal::MarshalError> {
if !self.allow_code {
return Err(self.remember_python_error(
self.vm
.new_value_error("unmarshalling code objects is disallowed"),
));
}
Ok(crate::builtins::PyCode::new_ref_with_bag(self.vm, code).into())
}
fn make_stop_iter(&self) -> Result<Self::Value, marshal::MarshalError> {
Ok(self.vm.ctx.exceptions.stop_iteration.to_owned().into())
}
fn make_list(
&self,
it: impl Iterator<Item = Self::Value>,
) -> Result<Self::Value, marshal::MarshalError> {
Ok(self.vm.ctx.new_list(it.collect()).into())
}
fn make_list_placeholder(
&self,
len: usize,
) -> Result<Option<Self::Value>, marshal::MarshalError> {
let elements = self.placeholder_elements(len)?;
Ok(Some(self.vm.ctx.new_list(elements).into()))
}
fn set_list_item(
&self,
list: &Self::Value,
index: usize,
value: Self::Value,
) -> Result<(), marshal::MarshalError> {
let list = list
.downcast_ref::<PyList>()
.ok_or(marshal::MarshalError::BadType)?;
list.borrow_vec_mut()[index] = value;
Ok(())
}
fn make_set(
&self,
it: impl Iterator<Item = Self::Value>,
) -> Result<Self::Value, marshal::MarshalError> {
let set = PySet::default().into_ref(&self.vm.ctx);
for elem in it {
set.add(elem, self.vm)
.map_err(|error| self.remember_python_error(error))?;
}
Ok(set.into())
}
fn make_set_placeholder(&self) -> Option<Self::Value> {
Some(PySet::default().into_ref(&self.vm.ctx).into())
}
fn insert_set_item(
&self,
set: &Self::Value,
value: Self::Value,
) -> Result<(), marshal::MarshalError> {
let set = set
.downcast_ref::<PySet>()
.ok_or(marshal::MarshalError::BadType)?;
set.add(value, self.vm)
.map_err(|error| self.remember_python_error(error))
}
fn make_frozenset(
&self,
it: impl Iterator<Item = Self::Value>,
) -> Result<Self::Value, marshal::MarshalError> {
PyFrozenSet::from_iter(self.vm, it)
.map(|set| set.to_pyobject(self.vm))
.map_err(|error| self.remember_python_error(error))
}
fn make_dict(
&self,
it: impl Iterator<Item = (Self::Value, Self::Value)>,
) -> Result<Self::Value, marshal::MarshalError> {
let dict = self.vm.ctx.new_dict();
for (k, v) in it {
dict.set_item(&*k, v, self.vm)
.map_err(|error| self.remember_python_error(error))?;
}
Ok(dict.into())
}
fn make_dict_placeholder(&self) -> Option<Self::Value> {
Some(self.vm.ctx.new_dict().into())
}
fn insert_dict_item(
&self,
dict: &Self::Value,
key: Self::Value,
value: Self::Value,
) -> Result<(), marshal::MarshalError> {
let dict = dict
.downcast_ref::<PyDict>()
.ok_or(marshal::MarshalError::BadType)?;
dict.set_item(&*key, value, self.vm)
.map_err(|error| self.remember_python_error(error))
}
fn make_slice(
&self,
start: Self::Value,
stop: Self::Value,
step: Self::Value,
) -> Result<Self::Value, marshal::MarshalError> {
use crate::builtins::PySlice;
let vm = self.vm;
Ok(PySlice {
start: if vm.is_none(&start) {
None
} else {
Some(start)
},
stop,
step: if vm.is_none(&step) { None } else { Some(step) },
}
.into_ref(&vm.ctx)
.into())
}
fn constant_bag(self) -> Self::ConstantBag {
PyVmBag(self.vm)
}
fn constant_ref_from_value(&self, value: &Self::Value) -> Option<Literal> {
Some(Literal::from(value.clone()))
}
fn bytes_from_value(&self, value: &Self::Value) -> Option<Vec<u8>> {
value
.downcast_ref::<PyBytes>()
.map(|bytes| bytes.as_bytes().to_vec())
}
fn str_from_value(&self, value: &Self::Value) -> Option<String> {
value
.downcast_ref::<PyStr>()
.map(|str| str.to_string_lossy().into_owned())
}
fn tuple_elements_from_value(&self, value: &Self::Value) -> Option<Vec<Self::Value>> {
value
.downcast_ref::<PyTuple>()
.map(|tuple| tuple.as_slice().to_vec())
}
}
fn deserialize_value(
rdr: &mut impl marshal::Read,
allow_code: bool,
vm: &VirtualMachine,
) -> PyResult<PyObjectRef> {
let pending_error = RefCell::new(None);
match marshal::deserialize_value(rdr, PyMarshalBag::new(vm, &pending_error, allow_code)) {
Ok(value) => Ok(value),
Err(error) => Err(pending_error.into_inner().unwrap_or_else(|| match error {
marshal::MarshalError::Eof => vm.new_eof_error("EOF read where not expected"),
marshal::MarshalError::EofObject => {
vm.new_eof_error("EOF read where object expected")
}
marshal::MarshalError::DataTooShort => vm.new_eof_error("marshal data too short"),
error @ marshal::MarshalError::NullObject => vm.new_type_error(error.to_string()),
error @ (marshal::MarshalError::BadSize(_)
| marshal::MarshalError::UnknownType
| marshal::MarshalError::InvalidRef) => {
vm.new_value_error(format!("bad marshal data ({error})"))
}
_ => vm.new_value_error("bad marshal data"),
})),
}
}
#[derive(FromArgs)]
struct LoadsArgs {
#[pyarg(positional)]
bytes: ArgBytesLike,
#[pyarg(named, default = true)]
allow_code: bool,
}
#[pyfunction]
fn loads(args: LoadsArgs, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
let LoadsArgs { bytes, allow_code } = args;
let buf = bytes.borrow_buf();
deserialize_value(&mut &buf[..], allow_code, vm)
}
#[derive(FromArgs)]
struct LoadArgs {
#[pyarg(positional)]
file: PyObjectRef,
#[pyarg(named, default = true)]
allow_code: bool,
}
#[pyfunction]
fn load(args: LoadArgs, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
let mut rdr = ReadableFile {
file: args.file,
vm,
buf: Vec::new(),
error: None,
};
let pending_error = RefCell::new(None);
let result = marshal::deserialize_value(
&mut rdr,
PyMarshalBag::new(vm, &pending_error, args.allow_code),
);
if let Some(err) = rdr.error.take() {
return Err(err);
}
match result {
Ok(value) => Ok(value),
Err(error) => Err(pending_error.into_inner().unwrap_or_else(|| match error {
marshal::MarshalError::Eof => vm.new_eof_error("EOF read where not expected"),
marshal::MarshalError::EofObject => {
vm.new_eof_error("EOF read where object expected")
}
marshal::MarshalError::DataTooShort => vm.new_eof_error("marshal data too short"),
error @ marshal::MarshalError::NullObject => vm.new_type_error(error.to_string()),
error @ (marshal::MarshalError::BadSize(_)
| marshal::MarshalError::UnknownType
| marshal::MarshalError::InvalidRef) => {
vm.new_value_error(format!("bad marshal data ({error})"))
}
_ => vm.new_value_error("bad marshal data"),
})),
}
}
struct ReadableFile<'a> {
file: PyObjectRef,
vm: &'a VirtualMachine,
buf: Vec<u8>,
error: Option<PyBaseExceptionRef>,
}
impl ReadableFile<'_> {
fn r_string(&mut self, n: usize) -> PyResult<()> {
self.buf.clear();
let bytearray = PyByteArray::from(vec![0u8; n]).into_ref(&self.vm.ctx);
let memoryview = PyMemoryView::from_object_with_flags(
bytearray.as_object(),
crate::protocol::BufferFlags::CONTIG,
self.vm,
)?
.into_ref(&self.vm.ctx);
let nread_obj = self.vm.call_method(&self.file, "readinto", (memoryview,))?;
let nread = nread_obj
.try_index(self.vm)?
.try_to_primitive::<isize>(self.vm)?;
let n_isize = isize::try_from(n).unwrap_or(isize::MAX);
if nread != n_isize {
if nread > n_isize {
return Err(self.vm.new_value_error(format!(
"read() returned too much data: {n} bytes requested, {nread} returned"
)));
}
return Err(self.vm.new_eof_error("EOF read where not expected"));
}
self.buf.extend_from_slice(&bytearray.borrow_buf());
Ok(())
}
}
impl marshal::Read for ReadableFile<'_> {
fn read_slice(&mut self, n: u32) -> Result<&[u8], marshal::MarshalError> {
if self.error.is_some() {
return Err(marshal::MarshalError::Eof);
}
if let Err(e) = self.r_string(n as usize) {
self.error = Some(e);
return Err(marshal::MarshalError::Eof);
}
Ok(&self.buf)
}
}
fn check_exact_type(obj: &PyObject, vm: &VirtualMachine) -> PyResult<()> {
let cls = obj.class();
if cls.is(PyBool::static_type()) {
return Ok(());
}
for base in [
PyInt::static_type(),
PyFloat::static_type(),
PyComplex::static_type(),
PyTuple::static_type(),
PyList::static_type(),
PyDict::static_type(),
PySet::static_type(),
PyFrozenSet::static_type(),
] {
if cls.fast_issubclass(base) && !cls.is(base) {
return Err(vm.new_value_error("unmarshallable object"));
}
}
Ok(())
}
}