use serde::de::Error as DeError;
use serde::de::IntoDeserializer;
use serde;
use super::SizeLimit;
use super::Error;
use super::Result;
use std::io::Read;
pub struct Deserializer<R, S: SizeLimit> {
reader: R,
store: u8,
shift: u8,
buffer: Vec<u8>,
size_limit: S,
}
macro_rules! de_uint {
($func:ident, $ty:ty) => {
fn $func(&mut self) -> Result<$ty> {
let v = if self.read_bit()? {
let h = ::std::mem::size_of::<$ty>();
let mut v = 0;
for x in 0 .. h {
v |= (self.read_byte()? as $ty) << x * 8;
if x != h - 1 && !self.read_bit()? {
break
}
}
v
} else {
0
};
Ok(v)
}
};
}
impl<'de, R: Read, S: SizeLimit> Deserializer<R, S> {
pub fn new(r: R, size_limit: S) -> Self {
Deserializer {
reader: r,
store: 0,
shift: 0,
size_limit: size_limit,
buffer: vec![],
}
}
fn read_bit(&mut self) -> Result<bool> {
self.probe_bits(1)?;
if self.shift == 0 {
let mut buf = [0; 1];
self.reader.read_exact(&mut buf)?;
self.store = buf[0];
}
let s = ((self.store >> self.shift) & 1) != 0;
self.shift = (self.shift + 1) % 8;
Ok(s)
}
fn read_byte(&mut self) -> Result<u8> {
self.probe_bits(8)?;
let mut buf = [0; 1];
self.reader.read_exact(&mut buf)?;
let s = buf[0];
let a = self.shift > 0;
let w = if a { (self.store >> self.shift) } else { 0 }
| (s << if a { (8 - self.shift) } else { 0 });
self.store = s;
Ok(w)
}
fn probe_bits(&mut self, count: u64) -> Result<()> {
self.size_limit.add(count)
}
de_uint!(de_u8, u8);
de_uint!(de_u16, u16);
de_uint!(de_u32, u32);
de_uint!(de_u64, u64);
}
impl<'de, 'a, R: Read, S> serde::Deserializer<'de>
for &'a mut Deserializer<R, S>
where
S: SizeLimit,
{
type Error = Error;
fn deserialize_any<V>(self, _visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
Err(Error::custom("zlo does not support Deserializer::deserialize_any"))
}
fn deserialize_bool<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_bool(self.read_bit()?)
}
fn deserialize_u8<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_u8(self.de_u8()?)
}
fn deserialize_u16<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_u16(self.de_u16()?)
}
fn deserialize_u32<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_u32(self.de_u32()?)
}
fn deserialize_u64<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_u64(self.de_u64()?)
}
fn deserialize_i8<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_i8(self.de_u8()? as i8)
}
fn deserialize_i16<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_i16(decode_zigzag_16(self.de_u16()?))
}
fn deserialize_i32<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_i32(decode_zigzag_32(self.de_u32()?))
}
fn deserialize_i64<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_i64(decode_zigzag_64(self.de_u64()?))
}
fn deserialize_f32<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
let sign = self.read_bit()?;
let exp = if self.read_bit()? {
self.read_byte()?
} else {
self.read_bit()? as u8 | ((self.read_bit()? as u8) << 1)
};
let frac = if self.read_bit()? {
((self.read_byte()? as u32) << 16)
| ((self.read_byte()? as u32) << 8)
| if self.read_bit()? { self.read_byte()? as u32 } else { 0 }
} else { 0 };
let bits = ((sign as u32) << 31) | ((exp as u32) << 23) | frac;
visitor.visit_f32(f32::from_bits(bits))
}
fn deserialize_f64<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
let sign = self.read_bit()?;
let exp = if self.read_bit()? {
((self.read_bit()? as u16) << 8)
| ((self.read_bit()? as u16) << 9)
| ((self.read_bit()? as u16) << 10)
| self.read_byte()? as u16
} else {
self.read_bit()? as u16 | ((self.read_bit()? as u16) << 1)
};
let frac = if self.read_bit()? {
((self.read_byte()? as u64) << 48)
| ((self.read_byte()? as u64) << 40)
| ((self.read_byte()? as u64) << 32)
| ((self.read_byte()? as u64) << 24)
| if self.read_bit()? {
((self.read_byte()? as u64) << 16)
| ((self.read_byte()? as u64) << 8)
| ((self.read_byte()? as u64) << 0)
} else { 0 }
} else { 0 };
let bits = ((sign as u64) << 63) | ((exp as u64) << 52) | frac;
visitor.visit_f64(f64::from_bits(bits))
}
fn deserialize_unit<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_unit()
}
fn deserialize_char<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
use std::str;
let error = || Error::InvalidEncoding {
desc: "Invalid char encoding",
detail: None,
}.into();
let mut buf = [0u8; 4];
let _ = self.reader.read_exact(&mut buf[..1]);
let width = utf8_char_width(buf[0]);
if width == 1 {
return visitor.visit_char(buf[0] as char);
}
if width == 0 {
return Err(error());
}
if self.reader.read_exact(&mut buf[1..width]).is_err() {
return Err(error());
}
let res = str::from_utf8(&buf[..width])
.ok()
.and_then(|s| s.chars().next())
.ok_or(error())?;
visitor.visit_char(res)
}
fn deserialize_str<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
let len: usize = serde::Deserialize::deserialize(&mut *self)?;
self.buffer.clear();
for _ in 0 .. len {
let b = self.read_byte()?;
self.buffer.push(b);
}
let s = ::std::str::from_utf8(&self.buffer[..])
.map_err(|err| Error::InvalidEncoding {
desc: "error while decoding utf8 string",
detail: Some(format!("Deserialize error: {}", err)),
})?;
visitor.visit_str(s)
}
fn deserialize_string<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_str(visitor)
}
fn deserialize_bytes<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
let len: usize = serde::Deserialize::deserialize(&mut *self)?;
self.buffer.clear();
for _ in 0 .. len {
let b = self.read_byte()?;
self.buffer.push(b);
}
visitor.visit_bytes(&self.buffer[..])
}
fn deserialize_byte_buf<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_bytes(visitor)
}
fn deserialize_enum<V>(
self,
_enum: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
impl<'de, 'a, R: 'a, S> serde::de::EnumAccess<'de>
for &'a mut Deserializer<R, S>
where
R: Read,
S: SizeLimit,
{
type Error = Error;
type Variant = Self;
fn variant_seed<V>(
self,
seed: V,
) -> Result<(V::Value, Self::Variant)>
where
V: serde::de::DeserializeSeed<'de>,
{
let idx: u32 = serde::de::Deserialize::deserialize(&mut *self)?;
let val: Result<_> = seed.deserialize(idx.into_deserializer());
Ok((val?, self))
}
}
visitor.visit_enum(self)
}
fn deserialize_tuple<V>(self, len: usize, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
struct Access<'a, R: Read + 'a, S: SizeLimit + 'a> {
deserializer: &'a mut Deserializer<R, S>,
len: usize,
}
impl<'de, 'a, 'b: 'a, R: Read + 'b, S: SizeLimit> serde::de::SeqAccess<'de>
for Access<'a, R, S>
{
type Error = Error;
fn next_element_seed<T>(
&mut self,
seed: T,
) -> Result<Option<T::Value>>
where
T: serde::de::DeserializeSeed<'de>,
{
if self.len > 0 {
self.len -= 1;
let value = serde::de::DeserializeSeed
::deserialize(seed, &mut *self.deserializer)?;
Ok(Some(value))
} else {
Ok(None)
}
}
fn size_hint(&self) -> Option<usize> {
Some(self.len)
}
}
visitor.visit_seq(Access {
deserializer: self,
len: len,
})
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
match self.read_bit()? {
false => visitor.visit_none(),
true => visitor.visit_some(&mut *self),
}
}
fn deserialize_seq<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
let len = serde::Deserialize::deserialize(&mut *self)?;
self.deserialize_tuple(len, visitor)
}
fn deserialize_map<V>(self, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
struct Access<'a, R: Read + 'a, S: SizeLimit + 'a> {
deserializer: &'a mut Deserializer<R, S>,
len: usize,
}
impl<'de, 'a, 'b: 'a, R: Read + 'b, S: SizeLimit> serde::de::MapAccess<'de>
for Access<'a, R, S>
{
type Error = Error;
fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>>
where
K: serde::de::DeserializeSeed<'de>,
{
if self.len > 0 {
self.len -= 1;
let key = serde::de::DeserializeSeed
::deserialize(seed, &mut *self.deserializer)?;
Ok(Some(key))
} else {
Ok(None)
}
}
fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value>
where
V: serde::de::DeserializeSeed<'de>,
{
let value = try!(serde::de::DeserializeSeed::deserialize(
seed,
&mut *self.deserializer,
));
Ok(value)
}
fn size_hint(&self) -> Option<usize> {
Some(self.len)
}
}
let len = serde::Deserialize::deserialize(&mut *self)?;
visitor.visit_map(Access {
deserializer: self,
len: len,
})
}
fn deserialize_struct<V>(
self,
_name: &str,
fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_tuple(fields.len(), visitor)
}
fn deserialize_identifier<V>(self, _visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
Err(Error::custom("zlo does not support Deserializer::deserialize_identifier"))
}
fn deserialize_newtype_struct<V>(
self,
_name: &str,
visitor: V,
) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_newtype_struct(self)
}
fn deserialize_unit_struct<V>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_unit()
}
fn deserialize_tuple_struct<V>(
self,
_name: &'static str,
len: usize,
visitor: V,
) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_tuple(len, visitor)
}
fn deserialize_ignored_any<V>(self, _visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
Err(Error::custom("zlo does not support Deserializer::deserialize_ignored_any"))
}
}
impl<'de, 'a, R: Read, S> serde::de::VariantAccess<'de>
for &'a mut Deserializer<R, S>
where
S: SizeLimit,
{
type Error = Error;
fn unit_variant(self) -> Result<()> {
Ok(())
}
fn newtype_variant_seed<T>(self, seed: T) -> Result<T::Value>
where
T: serde::de::DeserializeSeed<'de>,
{
serde::de::DeserializeSeed::deserialize(seed, self)
}
fn tuple_variant<V>(self, len: usize, visitor: V) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
serde::de::Deserializer::deserialize_tuple(self, len, visitor)
}
fn struct_variant<V>(
self,
fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value>
where
V: serde::de::Visitor<'de>,
{
serde::de::Deserializer::deserialize_tuple(self, fields.len(), visitor)
}
}
macro_rules! def_dec_zigzag {
($func:ident, $in:ty, $out:ty) => {
fn $func(v: $in) -> $out {
if (v % 2) == 1 {
(v / 2).wrapping_add(1).wrapping_neg() as $out
} else {
(v / 2) as $out
}
}
}
}
def_dec_zigzag!(decode_zigzag_16, u16, i16);
def_dec_zigzag!(decode_zigzag_32, u32, i32);
def_dec_zigzag!(decode_zigzag_64, u64, i64);
static UTF8_CHAR_WIDTH: [u8; 256] = [
1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,
1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1, 1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,
1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1, 1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,
1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1, 1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,
1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1, 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,
0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,
0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, 0,0,2,2,2,2,2,2,2,2,2,2,2,2,2,2,
2,2,2,2,2,2,2,2,2,2,2,2,2,2,2,2, 3,3,3,3,3,3,3,3,3,3,3,3,3,3,3,3, 4,4,4,4,4,0,0,0,0,0,0,0,0,0,0,0, ];
fn utf8_char_width(b: u8) -> usize {
UTF8_CHAR_WIDTH[b as usize] as usize
}