serde-beve 1.0.0

A BEVE data format for Serde
Documentation
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
    }
}