use super::Deserializer;
use crate::{error::Error, headers::ArrayKind};
use serde::{
de::{SeqAccess, Visitor},
forward_to_deserialize_any,
};
use std::io::Read;
pub struct SeqDeserializer<'a, R: Read> {
deserializer: &'a mut Deserializer<R>,
len: usize,
index: usize,
kind: ArrayKind,
}
impl<'a, R: Read> SeqDeserializer<'a, R> {
pub fn new(deserializer: &'a mut Deserializer<R>, len: usize, kind: ArrayKind) -> Self {
Self {
deserializer,
len,
kind,
index: 0,
}
}
}
impl<'a, 'de, R: Read> SeqAccess<'de> for SeqDeserializer<'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 {
return Ok(None);
}
let out = seed.deserialize(&mut *self)?;
self.index += 1;
Ok(Some(out))
}
fn size_hint(&self) -> Option<usize> {
Some(self.len)
}
}
macro_rules! deserialize_type {
($fn:ident, $kind:ident, $visitor:ident, $getter:ident) => {
fn $fn<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self.kind {
ArrayKind::$kind => visitor.$visitor(self.deserializer.$getter()?),
ArrayKind::Generic => self.deserializer.$fn(visitor),
found => Err(Error::MismatchedElementType {
expected: ArrayKind::$kind,
found,
}),
}
}
};
}
impl<'a, 'de, R: Read> serde::Deserializer<'de> for &mut SeqDeserializer<'a, R> {
type Error = Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self.kind {
ArrayKind::Generic => self.deserializer.deserialize_any(visitor),
ArrayKind::String => self.deserialize_string(visitor),
ArrayKind::Boolean => self.deserialize_bool(visitor),
ArrayKind::I8 => self.deserialize_i8(visitor),
ArrayKind::I16 => self.deserialize_i16(visitor),
ArrayKind::I32 => self.deserialize_i32(visitor),
ArrayKind::I64 => self.deserialize_i64(visitor),
ArrayKind::I128 => self.deserialize_i128(visitor),
ArrayKind::U8 => self.deserialize_u8(visitor),
ArrayKind::U16 => self.deserialize_u16(visitor),
ArrayKind::U32 => self.deserialize_u32(visitor),
ArrayKind::U64 => self.deserialize_u64(visitor),
ArrayKind::U128 => self.deserialize_u128(visitor),
ArrayKind::BF16 => visitor.visit_f32(self.deserializer.get_bf16_value()?),
ArrayKind::F16 => visitor.visit_f32(self.deserializer.get_f16_value()?),
ArrayKind::F32 => self.deserialize_f32(visitor),
ArrayKind::F64 => self.deserialize_f64(visitor),
ArrayKind::F128 | ArrayKind::Complex => unreachable!(),
}
}
deserialize_type!(deserialize_i8, I8, visit_i8, get_i8_value);
deserialize_type!(deserialize_i16, I16, visit_i16, get_i16_value);
deserialize_type!(deserialize_i32, I32, visit_i32, get_i32_value);
deserialize_type!(deserialize_i64, I64, visit_i64, get_i64_value);
deserialize_type!(deserialize_i128, I128, visit_i128, get_i128_value);
deserialize_type!(deserialize_u8, U8, visit_u8, get_u8_value);
deserialize_type!(deserialize_u16, U16, visit_u16, get_u16_value);
deserialize_type!(deserialize_u32, U32, visit_u32, get_u32_value);
deserialize_type!(deserialize_u64, U64, visit_u64, get_u64_value);
deserialize_type!(deserialize_u128, U128, visit_u128, get_u128_value);
deserialize_type!(deserialize_f32, F32, visit_f32, get_f32_value);
deserialize_type!(deserialize_f64, F64, visit_f64, get_f64_value);
deserialize_type!(deserialize_string, String, visit_string, get_string_value);
fn deserialize_bool<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self.kind {
ArrayKind::Boolean => {
let sub_index = self.index % 8;
let byte = if sub_index == 7 {
self.deserializer.get_byte()?
} else {
self.deserializer.peek_byte()?
};
let bit = byte & (1 << sub_index);
visitor.visit_bool(bit != 0)
}
ArrayKind::Generic => self.deserializer.deserialize_bool(visitor),
found => Err(Error::MismatchedElementType {
expected: ArrayKind::Boolean,
found,
}),
}
}
forward_to_deserialize_any! {
char str bytes byte_buf option unit unit_struct newtype_struct seq tuple tuple_struct map struct enum identifier ignored_any
}
}