use num_traits::ToPrimitive;
use crate::{
AsObject, PyObject, PyObjectRef, PyResult, VirtualMachine,
protocol::{BufferFlags, PyBuffer},
};
const BYTES_ELEMENT_ERROR: &str = "bytes must be in range(0, 256)";
const BYTE_ELEMENT_ERROR: &str = "byte must be in range(0, 256)";
pub fn bytes_from_object(vm: &VirtualMachine, obj: &PyObject) -> PyResult<Vec<u8>> {
collect_bytes(vm, obj, true, BYTES_ELEMENT_ERROR, |name| {
format!("cannot convert '{name}' object to bytes")
})
}
pub fn bytearray_from_object(vm: &VirtualMachine, obj: &PyObject) -> PyResult<Vec<u8>> {
collect_bytes(vm, obj, false, BYTE_ELEMENT_ERROR, |name| {
format!("cannot convert '{name}' object to bytearray")
})
}
pub fn bytearray_extend_from_object(vm: &VirtualMachine, obj: &PyObject) -> PyResult<Vec<u8>> {
collect_bytes(vm, obj, true, BYTE_ELEMENT_ERROR, |name| {
format!("can't extend bytearray with {name}")
})
}
fn collect_bytes(
vm: &VirtualMachine,
obj: &PyObject,
measured: bool,
element: &'static str,
unusable: impl FnOnce(&str) -> String,
) -> PyResult<Vec<u8>> {
if obj.check_buffer() {
let buffer = PyBuffer::from_object(vm, obj, BufferFlags::FULL_RO)?;
return Ok(buffer.contiguous_or_collect(|bytes| bytes.to_vec()));
}
if !obj.fast_isinstance(vm.ctx.types.str_type) {
let cls = obj.class();
if cls.slots().iter.load().is_none() && !cls.has_attr(identifier!(vm, __getitem__)) {
return Err(vm.new_type_error(unusable(&cls.name())));
}
let value = |x: PyObjectRef| value_from_object_with(vm, &x, element);
let elements = if measured {
vm.map_iterable_object_sized(obj, value)
} else {
vm.map_iterable_object(obj, value)
};
if let Ok(elements) = elements {
return elements;
}
}
Err(vm.new_type_error("can assign only bytes, buffers, or iterables of ints in range(0, 256)"))
}
pub fn value_from_object(vm: &VirtualMachine, obj: &PyObject) -> PyResult<u8> {
value_from_object_with(vm, obj, BYTE_ELEMENT_ERROR)
}
fn value_from_object_with(
vm: &VirtualMachine,
obj: &PyObject,
element: &'static str,
) -> PyResult<u8> {
obj.try_index(vm)?
.as_bigint()
.to_u8()
.ok_or_else(|| vm.new_value_error(element))
}