use crate::error::{Error, Result};
use serde::de::{
self, DeserializeSeed, EnumAccess, IntoDeserializer, MapAccess, SeqAccess, VariantAccess,
Visitor,
};
use serde::Deserialize;
use std::convert::TryFrom;
pub fn from_slice<'a, T: Deserialize<'a>>(input: &'a [u8]) -> Result<T> {
let mut deserializer = Deserializer { input };
let value = T::deserialize(&mut deserializer)?;
if deserializer.input.is_empty() {
Ok(value)
} else {
Err(Error::TrailingBytes(deserializer.input.len()))
}
}
pub struct Deserializer<'de> {
input: &'de [u8],
}
impl<'de> Deserializer<'de> {
#[inline]
fn take(&mut self, n: usize) -> Result<&'de [u8]> {
if self.input.len() < n {
return Err(Error::Eof);
}
let (head, tail) = self.input.split_at(n);
self.input = tail;
Ok(head)
}
#[inline]
fn take_array<const N: usize>(&mut self) -> Result<&'de [u8; N]> {
let (head, tail) = self.input.split_first_chunk::<N>().ok_or(Error::Eof)?;
self.input = tail;
Ok(head)
}
#[inline]
fn read_u8(&mut self) -> Result<u8> {
let (&first, tail) = self.input.split_first().ok_or(Error::Eof)?;
self.input = tail;
Ok(first)
}
#[inline]
fn read_len(&mut self) -> Result<usize> {
usize::try_from(self.read_u64()?).map_err(|_| Error::Eof)
}
#[inline]
fn read_u32(&mut self) -> Result<u32> {
Ok(u32::from_le_bytes(*self.take_array::<4>()?))
}
#[inline]
fn read_u64(&mut self) -> Result<u64> {
Ok(u64::from_le_bytes(*self.take_array::<8>()?))
}
}
macro_rules! read_le {
($self:ident, $ty:ty) => {
<$ty>::from_le_bytes(*$self.take_array::<{ std::mem::size_of::<$ty>() }>()?)
};
}
impl<'de> de::Deserializer<'de> for &mut Deserializer<'de> {
type Error = Error;
#[inline]
fn deserialize_any<V: Visitor<'de>>(self, _visitor: V) -> Result<V::Value> {
Err(Error::NotSupported)
}
#[inline]
fn deserialize_bool<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
match self.read_u8()? {
0 => visitor.visit_bool(false),
1 => visitor.visit_bool(true),
other => Err(Error::InvalidBool(other)),
}
}
#[inline]
fn deserialize_i8<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i8(read_le!(self, i8))
}
#[inline]
fn deserialize_i16<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i16(read_le!(self, i16))
}
#[inline]
fn deserialize_i32<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i32(read_le!(self, i32))
}
#[inline]
fn deserialize_i64<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i64(read_le!(self, i64))
}
#[inline]
fn deserialize_i128<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i128(read_le!(self, i128))
}
#[inline]
fn deserialize_u8<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u8(self.read_u8()?)
}
#[inline]
fn deserialize_u16<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u16(read_le!(self, u16))
}
#[inline]
fn deserialize_u32<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u32(self.read_u32()?)
}
#[inline]
fn deserialize_u64<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u64(self.read_u64()?)
}
#[inline]
fn deserialize_u128<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u128(read_le!(self, u128))
}
#[inline]
fn deserialize_f32<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_f32(read_le!(self, f32))
}
#[inline]
fn deserialize_f64<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_f64(read_le!(self, f64))
}
#[inline]
fn deserialize_char<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
let width = utf8_width(self.input.first().copied().ok_or(Error::Eof)?);
if width == 0 {
return Err(Error::InvalidChar);
}
let bytes = self.take(width)?;
let s = std::str::from_utf8(bytes).map_err(|_| Error::InvalidUtf8)?;
let mut chars = s.chars();
let c = chars.next().ok_or(Error::InvalidChar)?;
if chars.next().is_some() {
return Err(Error::InvalidChar);
}
visitor.visit_char(c)
}
#[inline]
fn deserialize_str<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
let len = self.read_len()?;
let bytes = self.take(len)?;
let s = std::str::from_utf8(bytes).map_err(|_| Error::InvalidUtf8)?;
visitor.visit_borrowed_str(s)
}
#[inline]
fn deserialize_string<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
self.deserialize_str(visitor)
}
#[inline]
fn deserialize_bytes<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
let len = self.read_len()?;
let bytes = self.take(len)?;
visitor.visit_borrowed_bytes(bytes)
}
#[inline]
fn deserialize_byte_buf<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
self.deserialize_bytes(visitor)
}
#[inline]
fn deserialize_option<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
match self.read_u8()? {
0 => visitor.visit_none(),
1 => visitor.visit_some(self),
other => Err(Error::InvalidOptionTag(other)),
}
}
#[inline]
fn deserialize_unit<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_unit()
}
#[inline]
fn deserialize_unit_struct<V: Visitor<'de>>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value> {
visitor.visit_unit()
}
#[inline]
fn deserialize_newtype_struct<V: Visitor<'de>>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value> {
visitor.visit_newtype_struct(self)
}
#[inline]
fn deserialize_seq<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
let len = self.read_len()?;
visitor.visit_seq(Counted::new(self, len))
}
#[inline]
fn deserialize_tuple<V: Visitor<'de>>(self, len: usize, visitor: V) -> Result<V::Value> {
visitor.visit_seq(Counted::new(self, len))
}
#[inline]
fn deserialize_tuple_struct<V: Visitor<'de>>(
self,
_name: &'static str,
len: usize,
visitor: V,
) -> Result<V::Value> {
visitor.visit_seq(Counted::new(self, len))
}
#[inline]
fn deserialize_map<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
let len = self.read_len()?;
visitor.visit_map(Counted::new(self, len))
}
#[inline]
fn deserialize_struct<V: Visitor<'de>>(
self,
_name: &'static str,
fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value> {
visitor.visit_seq(Counted::new(self, fields.len()))
}
#[inline]
fn deserialize_enum<V: Visitor<'de>>(
self,
_name: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value> {
visitor.visit_enum(self)
}
#[inline]
fn deserialize_identifier<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
self.deserialize_u32(visitor)
}
#[inline]
fn deserialize_ignored_any<V: Visitor<'de>>(self, _visitor: V) -> Result<V::Value> {
Err(Error::NotSupported)
}
}
struct Counted<'a, 'de: 'a> {
de: &'a mut Deserializer<'de>,
remaining: usize,
}
impl<'a, 'de> Counted<'a, 'de> {
#[inline]
fn new(de: &'a mut Deserializer<'de>, len: usize) -> Self {
Counted { de, remaining: len }
}
}
impl<'de, 'a> SeqAccess<'de> for Counted<'a, 'de> {
type Error = Error;
#[inline]
fn next_element_seed<T: DeserializeSeed<'de>>(&mut self, seed: T) -> Result<Option<T::Value>> {
if self.remaining == 0 {
return Ok(None);
}
self.remaining -= 1;
seed.deserialize(&mut *self.de).map(Some)
}
#[inline]
fn size_hint(&self) -> Option<usize> {
Some(self.remaining)
}
}
impl<'de, 'a> MapAccess<'de> for Counted<'a, 'de> {
type Error = Error;
#[inline]
fn next_key_seed<K: DeserializeSeed<'de>>(&mut self, seed: K) -> Result<Option<K::Value>> {
if self.remaining == 0 {
return Ok(None);
}
self.remaining -= 1;
seed.deserialize(&mut *self.de).map(Some)
}
#[inline]
fn next_value_seed<V: DeserializeSeed<'de>>(&mut self, seed: V) -> Result<V::Value> {
seed.deserialize(&mut *self.de)
}
#[inline]
fn size_hint(&self) -> Option<usize> {
Some(self.remaining)
}
}
impl<'de> EnumAccess<'de> for &mut Deserializer<'de> {
type Error = Error;
type Variant = Self;
#[inline]
fn variant_seed<V: DeserializeSeed<'de>>(self, seed: V) -> Result<(V::Value, Self::Variant)> {
let index = self.read_u32()?;
let variant_de: de::value::U32Deserializer<Error> = index.into_deserializer();
let value = seed.deserialize(variant_de)?;
Ok((value, self))
}
}
impl<'de> VariantAccess<'de> for &mut Deserializer<'de> {
type Error = Error;
#[inline]
fn unit_variant(self) -> Result<()> {
Ok(())
}
#[inline]
fn newtype_variant_seed<T: DeserializeSeed<'de>>(self, seed: T) -> Result<T::Value> {
seed.deserialize(self)
}
#[inline]
fn tuple_variant<V: Visitor<'de>>(self, len: usize, visitor: V) -> Result<V::Value> {
visitor.visit_seq(Counted::new(self, len))
}
#[inline]
fn struct_variant<V: Visitor<'de>>(
self,
fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value> {
visitor.visit_seq(Counted::new(self, fields.len()))
}
}
#[inline]
fn utf8_width(first: u8) -> usize {
match first {
0x00..=0x7F => 1,
0xC2..=0xDF => 2,
0xE0..=0xEF => 3,
0xF0..=0xF4 => 4,
_ => 0,
}
}