use std::io::Read;
use serde::{de::SeqAccess, forward_to_deserialize_any};
use crate::{ArrayKind, Error, de::seq::SeqDeserializer};
use super::Deserializer;
#[derive(Clone, Copy, PartialEq, Eq)]
pub enum ComplexKind {
I8,
I16,
I32,
I64,
I128,
U8,
U16,
U32,
U64,
U128,
F32,
F64,
}
impl From<ComplexKind> for ArrayKind {
fn from(kind: ComplexKind) -> Self {
match kind {
ComplexKind::I8 => ArrayKind::I8,
ComplexKind::I16 => ArrayKind::I16,
ComplexKind::I32 => ArrayKind::I32,
ComplexKind::I64 => ArrayKind::I64,
ComplexKind::I128 => ArrayKind::I128,
ComplexKind::U8 => ArrayKind::U8,
ComplexKind::U16 => ArrayKind::U16,
ComplexKind::U32 => ArrayKind::U32,
ComplexKind::U64 => ArrayKind::U64,
ComplexKind::U128 => ArrayKind::U128,
ComplexKind::F32 => ArrayKind::F32,
ComplexKind::F64 => ArrayKind::F64,
}
}
}
pub struct ComplexDeserializer<'a, R: Read> {
deserializer: &'a mut Deserializer<R>,
kind: ComplexKind,
index: usize,
}
impl<'a, R: Read> ComplexDeserializer<'a, R> {
pub fn new(deserializer: &'a mut Deserializer<R>, kind: ComplexKind) -> Self {
Self {
deserializer,
kind,
index: 0,
}
}
fn ensure_kind(&mut self, expected: ComplexKind) -> Result<(), Error> {
if self.kind == expected {
Ok(())
} else {
Err(Error::MismatchedElementType {
expected: expected.into(),
found: self.kind.into(),
})
}
}
}
impl<'a, 'de, R: Read> SeqAccess<'de> for ComplexDeserializer<'a, R> {
type Error = Error;
fn next_element_seed<T>(&mut self, seed: T) -> Result<Option<T::Value>, Self::Error>
where
T: serde::de::DeserializeSeed<'de>,
{
if self.index == 2 {
Ok(None)
} else {
self.index += 1;
seed.deserialize(self).map(Some)
}
}
fn size_hint(&self) -> Option<usize> {
Some(2)
}
}
macro_rules! deserialize_number {
($fn:ident, $kind:ident, $visitor:ident, $getter:ident) => {
fn $fn<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
self.ensure_kind(ComplexKind::$kind)?;
visitor.$visitor(self.deserializer.$getter()?)
}
};
}
impl<'a, 'de, R: Read> serde::Deserializer<'de> for &mut ComplexDeserializer<'a, R> {
type Error = Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let mut seq = SeqDeserializer::new(self.deserializer, 2, self.kind.into());
seq.deserialize_any(visitor)
}
deserialize_number!(deserialize_i8, I8, visit_i8, get_i8_value);
deserialize_number!(deserialize_i16, I16, visit_i16, get_i16_value);
deserialize_number!(deserialize_i32, I32, visit_i32, get_i32_value);
deserialize_number!(deserialize_i64, I64, visit_i64, get_i64_value);
deserialize_number!(deserialize_i128, I128, visit_i128, get_i128_value);
deserialize_number!(deserialize_u8, U8, visit_u8, get_u8_value);
deserialize_number!(deserialize_u16, U16, visit_u16, get_u16_value);
deserialize_number!(deserialize_u32, U32, visit_u32, get_u32_value);
deserialize_number!(deserialize_u64, U64, visit_u64, get_u64_value);
deserialize_number!(deserialize_u128, U128, visit_u128, get_u128_value);
deserialize_number!(deserialize_f32, F32, visit_f32, get_f32_value);
deserialize_number!(deserialize_f64, F64, visit_f64, get_f64_value);
forward_to_deserialize_any! {
bool char str string bytes byte_buf option unit unit_struct newtype_struct seq tuple tuple_struct map struct enum identifier ignored_any
}
}
pub struct ComplexArrayDeserializer<'a, R: Read> {
deserializer: &'a mut Deserializer<R>,
len: usize,
kind: ComplexKind,
index: usize,
}
impl<'a, R: Read> ComplexArrayDeserializer<'a, R> {
pub fn new(deserializer: &'a mut Deserializer<R>, len: usize, kind: ComplexKind) -> Self {
Self {
deserializer,
len,
kind,
index: 0,
}
}
}
impl<'a, 'de, R: Read> SeqAccess<'de> for ComplexArrayDeserializer<'a, R> {
type Error = Error;
fn next_element_seed<T>(&mut self, seed: T) -> Result<Option<T::Value>, Self::Error>
where
T: serde::de::DeserializeSeed<'de>,
{
if self.index == self.len {
Ok(None)
} else {
self.index += 1;
seed.deserialize(self).map(Some)
}
}
fn size_hint(&self) -> Option<usize> {
Some(self.len)
}
}
impl<'a, 'de, R: Read> serde::Deserializer<'de> for &mut ComplexArrayDeserializer<'a, R> {
type Error = Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_seq(visitor)
}
fn deserialize_seq<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_seq(ComplexDeserializer::new(self.deserializer, self.kind))
}
fn deserialize_tuple<V>(self, _len: usize, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_seq(ComplexDeserializer::new(self.deserializer, self.kind))
}
forward_to_deserialize_any! {
i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 bool char str string bytes byte_buf option unit unit_struct newtype_struct tuple_struct map struct enum identifier ignored_any
}
}