#![allow(clippy::cast_sign_loss)]
#![allow(clippy::cast_possible_wrap)]
#![allow(clippy::cast_lossless)]
use super::{ignored::Ignored, DeserializeSeed, Error, Kind, Result};
use crate::{tag::Tag, BignumRef, Deserialize, Fixnum, FromPrimitive, Sym, Visitor};
#[derive(Debug, Clone)]
pub struct Deserializer<'de> {
pub(crate) cursor: Cursor<'de>,
objtable: Vec<usize>,
stack: Vec<usize>,
is_reading_instance: bool,
sym_table: Vec<&'de Sym>,
}
#[derive(Debug, Clone)]
pub(crate) struct Cursor<'de> {
pub(crate) input: &'de [u8],
pub(crate) position: usize,
}
struct InstanceAccess<'de, 'a> {
deserializer: &'a mut Deserializer<'de>,
len: &'a mut usize,
index: &'a mut usize,
}
struct IvarAccess<'de, 'a> {
deserializer: &'a mut Deserializer<'de>,
len: usize,
index: &'a mut usize,
state: MapState,
}
struct ArrayAccess<'de, 'a> {
deserializer: &'a mut Deserializer<'de>,
len: usize,
index: &'a mut usize,
}
struct HashAccess<'de, 'a> {
deserializer: &'a mut Deserializer<'de>,
len: usize,
index: &'a mut usize,
state: MapState,
}
enum MapState {
Key,
Value,
}
impl<'de> Cursor<'de> {
fn new(input: &'de [u8]) -> Self {
Self { input, position: 0 }
}
fn seek(&mut self, position: usize) {
self.position = position;
}
fn peek_byte(&self) -> Result<u8> {
self.input
.get(self.position)
.copied()
.ok_or(Error { kind: Kind::Eof })
}
fn next_byte(&mut self) -> Result<u8> {
let byte = self.peek_byte()?;
self.position += 1;
Ok(byte)
}
fn peek_tag(&self) -> Result<Tag> {
let byte = self.peek_byte()?;
Tag::from_u8(byte).ok_or(Error {
kind: Kind::WrongTag(byte),
})
}
fn next_tag(&mut self) -> Result<Tag> {
let byte = self.next_byte()?;
Tag::from_u8(byte).ok_or(Error {
kind: Kind::WrongTag(byte),
})
}
fn next_bytes_dyn(&mut self, length: usize) -> Result<&'de [u8]> {
let new_position = self
.position
.checked_add(length)
.ok_or(Error { kind: Kind::Eof })?;
if new_position > self.input.len() {
return Err(Error { kind: Kind::Eof });
}
let ret = &self.input[self.position..new_position];
self.position = new_position;
Ok(ret)
}
}
impl<'de> Deserializer<'de> {
pub fn new(input: &'de [u8]) -> Result<Self> {
let mut cursor = Cursor::new(input);
if input.len() < 2 {
return Err(Error { kind: Kind::Eof });
}
let v1 = cursor.next_byte()?;
let v2 = cursor.next_byte()?;
if [v1, v2] != [4, 8] {
return Err(Error {
kind: Kind::VersionError([v1, v2]),
});
}
Ok(Self {
cursor,
objtable: vec![],
sym_table: vec![],
is_reading_instance: false,
stack: vec![],
})
}
pub fn deserialize_value<T>(&mut self) -> Result<T>
where
T: Deserialize<'de>,
{
T::deserialize(self)
}
pub fn current_position(&self) -> usize {
self.cursor.position
}
pub fn data(&self) -> &'de [u8] {
self.cursor.input
}
fn read_fixnum(&mut self) -> Result<Fixnum> {
let c = self.cursor.next_byte()? as i8;
Ok(match c {
0 => 0i16.into(),
5..=127 => (c - 5).into(),
-128..=-5 => (c + 5).into(),
1..=4 => {
let mut x = 0;
for i in 0..c {
let n = self.cursor.next_byte()? as i32;
let n = n << (8 * i);
x |= n;
}
FromPrimitive::from_i32(x).unwrap()
}
-4..=-1 => {
let mut x = -1;
for i in 0..-c {
let a = !(0xFF << (8 * i)); let b = self.cursor.next_byte()? as i32;
let b = b << (8 * i);
x = (x & a) | b;
}
FromPrimitive::from_i32(x).unwrap()
}
})
}
fn read_usize(&mut self) -> Result<usize> {
num_traits::ToPrimitive::to_usize(&self.read_fixnum()?).ok_or(Error { kind: Kind::Eof })
}
#[allow(clippy::panic_in_result_fn)]
fn read_float(&mut self) -> Result<f64> {
let out = self.read_bytes_len()?;
if let Some(terminator_idx) = out.iter().position(|v| *v == 0) {
let (str, [0, mantissa @ ..]) = out.split_at(terminator_idx) else {
unreachable!();
};
let float = str::parse::<f64>(&String::from_utf8_lossy(str)).map_err(|err| Error {
kind: Kind::Message(err.to_string()),
})?;
let transmuted = u64::from_ne_bytes(float.to_ne_bytes());
if mantissa.len() > 4 {
return Err(Error {
kind: Kind::ParseFloatMantissaTooLong,
});
}
let (mantissa, mask) = mantissa.iter().fold((0u64, 0u64), |(acc, mask), v| {
((acc << 8) | u64::from(*v), (mask << 8) | 0xFF)
});
let transmuted = (transmuted & !mask) | mantissa;
Ok(f64::from_ne_bytes(transmuted.to_ne_bytes()))
} else {
Ok(
str::parse::<f64>(&String::from_utf8_lossy(out)).map_err(|err| Error {
kind: Kind::Message(err.to_string()),
})?,
)
}
}
fn read_symbol(&mut self) -> Result<&'de Sym> {
let out = self.read_str_len()?;
let sym = Sym::new(out);
if self.stack.is_empty() {
self.sym_table.push(sym);
}
Ok(sym)
}
fn read_symlink(&mut self) -> Result<&'de Sym> {
let index = self.read_usize()?;
self.sym_table.get(index).copied().ok_or(Error {
kind: Kind::UnresolvedSymlink(index),
})
}
fn read_symbol_either(&mut self) -> Result<&'de Sym> {
match self.cursor.next_tag()? {
Tag::Symbol => self.read_symbol(),
Tag::Symlink => self.read_symlink(),
t => Err(Error {
kind: Kind::ExpectedSymbol(t),
}),
}
}
fn register_obj(&mut self) {
if !self.stack.is_empty() || self.is_reading_instance {
self.is_reading_instance = false; return;
}
self.objtable.push(self.cursor.position);
}
fn read_bytes_len(&mut self) -> Result<&'de [u8]> {
let len = self.read_usize()?;
self.cursor.next_bytes_dyn(len)
}
fn read_str_len(&mut self) -> Result<&'de str> {
let len = self.read_usize()?;
let bytes = self.cursor.next_bytes_dyn(len)?;
std::str::from_utf8(bytes).map_err(|e| Error {
kind: Kind::SymbolInvalidUTF8(e),
})
}
}
impl<'de> super::DeserializerTrait<'de> for &mut Deserializer<'de> {
#[allow(clippy::too_many_lines)]
fn deserialize<V>(self, visitor: V) -> Result<V::Value>
where
V: Visitor<'de>,
{
if self.cursor.peek_tag()?.is_object_link_referenceable() {
self.register_obj();
}
match self.cursor.next_tag()? {
Tag::Nil => visitor.visit_nil(),
Tag::True => visitor.visit_bool(true),
Tag::False => visitor.visit_bool(false),
Tag::Fixnum => visitor.visit_fixnum(self.read_fixnum()?),
Tag::Float => visitor.visit_f64(self.read_float()?),
Tag::Bignum => {
let is_negative = self.cursor.next_byte()? == b'-';
let len = self
.read_usize()?
.checked_mul(2)
.ok_or(Error { kind: Kind::Eof })?;
let le_bytes = self.cursor.next_bytes_dyn(len)?;
if let Some(bignum) = BignumRef::from_le_bytes(is_negative, le_bytes) {
visitor.visit_bignum(bignum)
} else {
visitor.visit_fixnum(Fixnum::from_le_bytes(is_negative, le_bytes).unwrap())
}
}
Tag::String => {
let data = self.read_bytes_len()?;
visitor.visit_string(data)
}
Tag::Array => {
let len = self.read_usize()?;
let mut index = 0;
let result = visitor.visit_array(ArrayAccess {
deserializer: self,
len,
index: &mut index,
})?;
while index < len {
index += 1;
Ignored::deserialize(&mut *self)?;
}
Ok(result)
}
Tag::Hash => {
let len = self.read_usize()?;
let mut index = 0;
let result = visitor.visit_hash(HashAccess {
deserializer: self,
len,
index: &mut index,
state: MapState::Value, })?;
while index < len {
index += 1;
Ignored::deserialize(&mut *self)?;
Ignored::deserialize(&mut *self)?;
}
Ok(result)
}
Tag::Symbol => visitor.visit_symbol(self.read_symbol()?),
Tag::Symlink => visitor.visit_symbol(self.read_symlink()?),
Tag::Instance => {
self.is_reading_instance = true;
let mut len = 0;
let mut index = 0;
let result = visitor.visit_instance(&mut InstanceAccess {
deserializer: &mut *self,
len: &mut len,
index: &mut index,
})?;
while index < len {
index += 1;
self.read_symbol_either()?;
Ignored::deserialize(&mut *self)?;
}
Ok(result)
}
Tag::Object => {
let class = self.read_symbol_either()?;
let len = self.read_usize()?;
let mut index = 0;
let result = visitor.visit_object(
class,
IvarAccess {
deserializer: self,
len,
index: &mut index,
state: MapState::Value, },
)?;
while index < len {
index += 1;
self.read_symbol_either()?;
Ignored::deserialize(&mut *self)?;
}
Ok(result)
}
Tag::ObjectLink => {
let index = self.read_usize()?;
let jump_target = self.objtable.get(index).copied().ok_or(Error {
kind: Kind::UnresolvedObjectlink(index),
})?;
if self.stack.contains(&self.cursor.position) {
return Err(Error {
kind: Kind::CircularReference,
});
}
self.stack.push(self.cursor.position);
self.cursor.seek(jump_target);
let result = self.deserialize(visitor);
self.cursor
.seek(self.stack.pop().expect("stack should not empty"));
result
}
Tag::UserDef => {
let class = self.read_symbol_either()?;
let data = self.read_bytes_len()?;
visitor.visit_user_data(class, data)
}
Tag::HashDefault => {
let len = self.read_usize()?;
let mut index = 0;
let result = visitor.visit_hash(HashAccess {
deserializer: self,
len,
index: &mut index,
state: MapState::Value, });
while index < len {
index += 1;
Ignored::deserialize(&mut *self)?;
Ignored::deserialize(&mut *self)?;
}
Ignored::deserialize(&mut *self)?;
result
}
Tag::UserClass => {
let class = self.read_symbol_either()?;
visitor.visit_user_class(class, &mut *self)
}
Tag::RawRegexp => {
let regex = self.read_bytes_len()?;
let flags = self.cursor.next_byte()?;
visitor.visit_regular_expression(regex, flags)
}
Tag::ClassRef => {
let class = self.read_str_len()?;
visitor.visit_class(Sym::new(class))
}
Tag::ModuleRef => {
let module = self.read_str_len()?;
visitor.visit_module(Sym::new(module))
}
Tag::Extended => {
let module = self.read_symbol_either()?;
visitor.visit_extended(module, &mut *self)
}
Tag::UserMarshal => {
let class = self.read_symbol_either()?;
visitor.visit_user_marshal(class, &mut *self)
}
Tag::Struct => {
let name = self.read_symbol_either()?;
let len = self.read_usize()?;
let mut index = 0;
let result = visitor.visit_struct(
name,
IvarAccess {
deserializer: self,
len,
index: &mut index,
state: MapState::Value, },
)?;
while index < len {
index += 1;
self.read_symbol_either()?;
Ignored::deserialize(&mut *self)?;
}
Ok(result)
}
Tag::Data => {
let class = self.read_symbol_either()?;
visitor.visit_data(class, &mut *self)
}
}
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value>
where
V: super::traits::VisitorOption<'de>,
{
if self.cursor.peek_tag()? == Tag::Nil {
self.cursor.next_byte()?;
visitor.visit_none()
} else {
visitor.visit_some(self)
}
}
fn deserialize_instance<V>(self, visitor: V) -> Result<V::Value>
where
V: super::traits::VisitorInstance<'de>,
{
if self.cursor.peek_tag()? == Tag::Instance {
self.register_obj(); self.is_reading_instance = true;
self.cursor.next_byte()?;
let mut len = 0;
let mut index = 0;
let result = visitor.visit_instance(&mut InstanceAccess {
deserializer: &mut *self,
len: &mut len,
index: &mut index,
})?;
while index < len {
index += 1;
self.read_symbol_either()?;
Ignored::deserialize(&mut *self)?;
}
Ok(result)
} else {
visitor.visit(self)
}
}
}
impl<'de, 'a> super::InstanceAccess<'de> for &'a mut InstanceAccess<'de, 'a> {
type IvarAccess = IvarAccess<'de, 'a>;
fn value_seed<V>(self, seed: V) -> Result<(V::Value, Self::IvarAccess)>
where
V: DeserializeSeed<'de>,
{
let result = seed.deserialize(&mut *self.deserializer)?;
let len = self.deserializer.read_usize()?;
*self.len = len;
Ok((
result,
IvarAccess {
deserializer: &mut *self.deserializer,
len,
index: self.index,
state: MapState::Value, },
))
}
}
impl<'de> super::IvarAccess<'de> for IvarAccess<'de, '_> {
fn next_ivar(&mut self) -> Result<Option<&'de Sym>> {
if *self.index >= self.len {
return Ok(None);
}
match self.state {
MapState::Key => {
return Err(Error {
kind: Kind::KeyAfterKey,
})
}
MapState::Value => self.state = MapState::Key,
}
*self.index += 1;
self.deserializer.read_symbol_either().map(Some)
}
fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value>
where
V: DeserializeSeed<'de>,
{
match self.state {
MapState::Value => {
return Err(Error {
kind: Kind::ValueAfterValue,
})
}
MapState::Key => self.state = MapState::Value,
}
seed.deserialize(&mut *self.deserializer)
}
fn len(&self) -> usize {
self.len
}
fn index(&self) -> usize {
*self.index
}
}
impl<'de> super::ArrayAccess<'de> for ArrayAccess<'de, '_> {
fn next_element_seed<T>(&mut self, seed: T) -> Result<Option<T::Value>>
where
T: DeserializeSeed<'de>,
{
if *self.index >= self.len {
return Ok(None);
}
*self.index += 1;
seed.deserialize(&mut *self.deserializer).map(Some)
}
fn len(&self) -> usize {
self.len
}
fn index(&self) -> usize {
*self.index
}
}
impl<'de> super::HashAccess<'de> for HashAccess<'de, '_> {
fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>>
where
K: DeserializeSeed<'de>,
{
if *self.index >= self.len {
return Ok(None);
}
match self.state {
MapState::Key => {
return Err(Error {
kind: Kind::KeyAfterKey,
})
}
MapState::Value => self.state = MapState::Key,
}
*self.index += 1;
seed.deserialize(&mut *self.deserializer).map(Some)
}
fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value>
where
V: DeserializeSeed<'de>,
{
match self.state {
MapState::Value => {
return Err(Error {
kind: Kind::ValueAfterValue,
})
}
MapState::Key => self.state = MapState::Value,
}
seed.deserialize(&mut *self.deserializer)
}
fn len(&self) -> usize {
self.len
}
fn index(&self) -> usize {
*self.index
}
}