mod complex;
mod enums;
mod map;
mod matrix;
mod seq;
use crate::{Error, error::SpecialType, headers::*};
use complex::{ComplexArrayDeserializer, ComplexDeserializer, ComplexKind};
use enums::EnumDeserializer;
use map::MapDeserializer;
use matrix::MatrixDeserializer;
use seq::SeqDeserializer;
use serde::{de::Visitor, forward_to_deserialize_any};
use std::io::Read;
pub struct Deserializer<R: Read> {
reader: R,
peek: Option<u8>,
}
macro_rules! deserialize_object {
($fn:ident, $header:ident, $kind:ident) => {
fn $fn<'de, V: Visitor<'de>>(
&mut self,
visitor: V,
) -> Result<V::Value, <&mut Self as serde::Deserializer<'_>>::Error> {
if self.get_byte()? != $header {
return Err(Error::WrongType {
expected: header_name($header),
found: header_name(self.get_byte()?),
});
}
let size = self.get_size()?;
visitor.visit_map(MapDeserializer::new(self, size, ObjectKind::$kind))
}
};
}
macro_rules! deserialize_array {
($fn:ident, $header:ident, $kind:ident) => {
fn $fn<'de, V: Visitor<'de>>(
&mut self,
visitor: V,
) -> Result<V::Value, <&mut Self as serde::Deserializer<'_>>::Error> {
if self.get_byte()? != $header {
return Err(Error::WrongType {
expected: header_name($header),
found: header_name(self.get_byte()?),
});
}
let size = self.get_size()?;
visitor.visit_seq(SeqDeserializer::new(self, size, ArrayKind::$kind))
}
};
}
impl<R: Read> Deserializer<R> {
pub fn new(reader: R) -> Self {
Self { reader, peek: None }
}
fn read_byte(&mut self) -> Result<u8, Error> {
let mut buf = [0];
self.reader.read_exact(&mut buf)?;
Ok(buf[0])
}
pub(self) fn get_byte(&mut self) -> Result<u8, Error> {
if let Some(peek) = self.peek.take() {
return Ok(peek);
}
self.read_byte()
}
pub(self) fn peek_byte(&mut self) -> Result<u8, Error> {
if let Some(peek) = self.peek {
return Ok(peek);
}
let read = self.read_byte()?;
self.peek = Some(read);
Ok(read)
}
fn get_u8_array(&mut self) -> Result<Vec<u8>, Error> {
let size = self.get_size()?;
let mut bytes = vec![0; size];
self.reader.read_exact(&mut bytes)?;
Ok(bytes)
}
fn deserialize_bf16<'de, V: Visitor<'de>>(&mut self, visitor: V) -> Result<V::Value, Error> {
if self.get_byte()? != BF16 {
return Err(Error::WrongType {
expected: header_name(BF16),
found: header_name(self.get_byte()?),
});
}
visitor.visit_f32(self.get_bf16_value()?)
}
fn deserialize_f16<'de, V: Visitor<'de>>(&mut self, visitor: V) -> Result<V::Value, Error> {
if self.get_byte()? != F16 {
return Err(Error::WrongType {
expected: header_name(F16),
found: header_name(self.get_byte()?),
});
}
visitor.visit_f32(self.get_f16_value()?)
}
deserialize_object!(deserialize_string_object, STRING_OBJECT, String);
deserialize_object!(deserialize_i8_object, I8_OBJECT, I8);
deserialize_object!(deserialize_i16_object, I16_OBJECT, I16);
deserialize_object!(deserialize_i32_object, I32_OBJECT, I32);
deserialize_object!(deserialize_i64_object, I64_OBJECT, I64);
deserialize_object!(deserialize_i128_object, I128_OBJECT, I128);
deserialize_object!(deserialize_u8_object, U8_OBJECT, U8);
deserialize_object!(deserialize_u16_object, U16_OBJECT, U16);
deserialize_object!(deserialize_u32_object, U32_OBJECT, U32);
deserialize_object!(deserialize_u64_object, U64_OBJECT, U64);
deserialize_object!(deserialize_u128_object, U128_OBJECT, U128);
deserialize_array!(deserialize_bf16_array, BF16_ARRAY, BF16);
deserialize_array!(deserialize_f16_array, F16_ARRAY, F16);
deserialize_array!(deserialize_f32_array, F32_ARRAY, F32);
deserialize_array!(deserialize_f64_array, F64_ARRAY, F64);
deserialize_array!(deserialize_i8_array, I8_ARRAY, I8);
deserialize_array!(deserialize_i16_array, I16_ARRAY, I16);
deserialize_array!(deserialize_i32_array, I32_ARRAY, I32);
deserialize_array!(deserialize_i64_array, I64_ARRAY, I64);
deserialize_array!(deserialize_i128_array, I128_ARRAY, I128);
deserialize_array!(deserialize_u8_array, U8_ARRAY, U8);
deserialize_array!(deserialize_u16_array, U16_ARRAY, U16);
deserialize_array!(deserialize_u32_array, U32_ARRAY, U32);
deserialize_array!(deserialize_u64_array, U64_ARRAY, U64);
deserialize_array!(deserialize_u128_array, U128_ARRAY, U128);
deserialize_array!(deserialize_bool_array, BOOL_ARRAY, Boolean);
deserialize_array!(deserialize_string_array, STRING_ARRAY, String);
pub(self) fn get_size(&mut self) -> Result<usize, Error> {
let first = self.get_byte()?;
let n_bytes = 2_usize.pow((first & 0b11) as u32);
if n_bytes == 1 {
return Ok((first as usize) >> 2);
}
let mut rest = vec![0; n_bytes - 1];
self.reader.read_exact(&mut rest)?;
#[cfg(target_pointer_width = "64")]
let mut bytes = [0; 8];
#[cfg(target_pointer_width = "32")]
let mut bytes = [0; 4];
#[cfg(target_pointer_width = "16")]
let mut bytes = [0; 2];
if rest.len() >= bytes.len() {
return Err(Error::TooLong);
}
bytes[0] = first;
for (i, byte) in rest.into_iter().enumerate() {
bytes[i + 1] = byte;
}
Ok(usize::from_le_bytes(bytes) >> 2)
}
pub(self) fn get_string_value(&mut self) -> Result<String, Error> {
let size = self.get_size()?;
let mut bytes = vec![0; size];
self.reader.read_exact(&mut bytes)?;
Ok(String::from_utf8(bytes)?)
}
pub(self) fn get_u8_value(&mut self) -> Result<u8, Error> {
self.get_byte()
}
fn get_num_value<T, const N: usize>(&mut self, f: fn([u8; N]) -> T) -> Result<T, Error> {
let mut bytes = [0; N];
self.reader.read_exact(&mut bytes)?;
Ok(f(bytes))
}
pub(self) fn get_u16_value(&mut self) -> Result<u16, Error> {
self.get_num_value(u16::from_le_bytes)
}
pub(self) fn get_u32_value(&mut self) -> Result<u32, Error> {
self.get_num_value(u32::from_le_bytes)
}
pub(self) fn get_u64_value(&mut self) -> Result<u64, Error> {
self.get_num_value(u64::from_le_bytes)
}
pub(self) fn get_u128_value(&mut self) -> Result<u128, Error> {
self.get_num_value(u128::from_le_bytes)
}
pub(self) fn get_i8_value(&mut self) -> Result<i8, Error> {
self.get_num_value(i8::from_le_bytes)
}
pub(self) fn get_i16_value(&mut self) -> Result<i16, Error> {
self.get_num_value(i16::from_le_bytes)
}
pub(self) fn get_i32_value(&mut self) -> Result<i32, Error> {
self.get_num_value(i32::from_le_bytes)
}
pub(self) fn get_i64_value(&mut self) -> Result<i64, Error> {
self.get_num_value(i64::from_le_bytes)
}
pub(self) fn get_i128_value(&mut self) -> Result<i128, Error> {
self.get_num_value(i128::from_le_bytes)
}
pub(self) fn get_bf16_value(&mut self) -> Result<f32, Error> {
#[cfg(feature = "half")]
{
self.get_num_value(half::bf16::from_le_bytes)
.map(|v| v.to_f32())
}
#[cfg(not(feature = "half"))]
{
Err(Error::UnsupportedDataType(SpecialType::BrainFloat))
}
}
pub(self) fn get_f16_value(&mut self) -> Result<f32, Error> {
#[cfg(feature = "half")]
{
self.get_num_value(half::f16::from_le_bytes)
.map(|v| v.to_f32())
}
#[cfg(not(feature = "half"))]
{
Err(Error::UnsupportedDataType(SpecialType::HalfFloat))
}
}
pub(self) fn get_f32_value(&mut self) -> Result<f32, Error> {
self.get_num_value(f32::from_le_bytes)
}
pub(self) fn get_f64_value(&mut self) -> Result<f64, Error> {
self.get_num_value(f64::from_le_bytes)
}
fn deserialize_complex<'de, V: Visitor<'de>>(&mut self, visitor: V) -> Result<V::Value, Error> {
match self.get_byte()? {
COMPLEX => {}
header => {
return Err(Error::WrongType {
expected: header_name(COMPLEX),
found: header_name(header),
});
}
}
let complex_header = self.get_byte()?;
let array = complex_header & 1 == 1;
let kind = match complex_header & NUM_TYPE_MASK {
I8_HEADER => ComplexKind::I8,
I16_HEADER => ComplexKind::I16,
I32_HEADER => ComplexKind::I32,
I64_HEADER => ComplexKind::I64,
I128_HEADER => ComplexKind::I128,
U8_HEADER => ComplexKind::U8,
U16_HEADER => ComplexKind::U16,
U32_HEADER => ComplexKind::U32,
U64_HEADER => ComplexKind::U64,
U128_HEADER => ComplexKind::U128,
F32_HEADER => ComplexKind::F32,
F64_HEADER => ComplexKind::F64,
_ => {
return Err(Error::InvalidComplexHeader);
}
};
if array {
let size = self.get_size()?;
visitor.visit_seq(ComplexArrayDeserializer::new(self, size, kind))
} else {
visitor.visit_seq(ComplexDeserializer::new(self, kind))
}
}
fn deserialize_matrix<'de, V: Visitor<'de>>(&mut self, visitor: V) -> Result<V::Value, Error> {
match self.get_byte()? {
MATRIX => {}
header => {
return Err(Error::WrongType {
expected: header_name(MATRIX),
found: header_name(header),
});
}
}
let layout = if self.get_byte()? & 1 == 1 {
"layout_right"
} else {
"layout_left"
};
visitor.visit_map(MatrixDeserializer::new(self, layout.to_string()))
}
}
macro_rules! deserialize_primitive {
($fn:ident, $header:ident, $visitor:ident, $getter:ident) => {
fn $fn<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self.get_byte()? {
$header => visitor.$visitor(self.$getter()?),
header => Err(Error::WrongType {
expected: header_name($header),
found: header_name(header),
}),
}
}
};
}
impl<'de, R: Read> serde::Deserializer<'de> for &mut Deserializer<R> {
type Error = Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
match self.peek_byte()? {
NULL => self.deserialize_unit(visitor),
TRUE | FALSE => self.deserialize_bool(visitor),
BF16 => self.deserialize_bf16(visitor),
F16 => self.deserialize_f16(visitor),
F32 => self.deserialize_f32(visitor),
F64 => self.deserialize_f64(visitor),
F128 => Err(Error::UnsupportedDataType(SpecialType::F128)),
I8 => self.deserialize_i8(visitor),
I16 => self.deserialize_i16(visitor),
I32 => self.deserialize_i32(visitor),
I64 => self.deserialize_i64(visitor),
I128 => self.deserialize_i128(visitor),
U8 => self.deserialize_u8(visitor),
U16 => self.deserialize_u16(visitor),
U32 => self.deserialize_u32(visitor),
U64 => self.deserialize_u64(visitor),
U128 => self.deserialize_u128(visitor),
STRING => self.deserialize_string(visitor),
STRING_OBJECT => self.deserialize_string_object(visitor),
I8_OBJECT => self.deserialize_i8_object(visitor),
I16_OBJECT => self.deserialize_i16_object(visitor),
I32_OBJECT => self.deserialize_i32_object(visitor),
I64_OBJECT => self.deserialize_i64_object(visitor),
I128_OBJECT => self.deserialize_i128_object(visitor),
U8_OBJECT => self.deserialize_u8_object(visitor),
U16_OBJECT => self.deserialize_u16_object(visitor),
U32_OBJECT => self.deserialize_u32_object(visitor),
U64_OBJECT => self.deserialize_u64_object(visitor),
U128_OBJECT => self.deserialize_u128_object(visitor),
BF16_ARRAY => self.deserialize_bf16_array(visitor),
F16_ARRAY => self.deserialize_f16_array(visitor),
F32_ARRAY => self.deserialize_f32_array(visitor),
F64_ARRAY => self.deserialize_f64_array(visitor),
F128_ARRAY => Err(Error::UnsupportedDataType(SpecialType::F128)),
I8_ARRAY => self.deserialize_i8_array(visitor),
I16_ARRAY => self.deserialize_i16_array(visitor),
I32_ARRAY => self.deserialize_i32_array(visitor),
I64_ARRAY => self.deserialize_i64_array(visitor),
I128_ARRAY => self.deserialize_i128_array(visitor),
U8_ARRAY => self.deserialize_u8_array(visitor),
U16_ARRAY => self.deserialize_u16_array(visitor),
U32_ARRAY => self.deserialize_u32_array(visitor),
U64_ARRAY => self.deserialize_u64_array(visitor),
U128_ARRAY => self.deserialize_u128_array(visitor),
BOOL_ARRAY => self.deserialize_bool_array(visitor),
STRING_ARRAY => self.deserialize_string_array(visitor),
GENERIC_ARRAY => self.deserialize_seq(visitor),
DELIMITER => {
self.get_byte()?;
visitor.visit_unit()
}
TAG => self.deserialize_enum("", &[], visitor),
MATRIX => self.deserialize_matrix(visitor),
COMPLEX => self.deserialize_complex(visitor),
RESERVED => Err(Error::Reserved),
header => Err(Error::InvalidHeader(header)),
}
}
fn deserialize_bool<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
match self.get_byte()? {
TRUE => visitor.visit_bool(true),
FALSE => visitor.visit_bool(false),
header => Err(Error::WrongType {
expected: header_name(TRUE),
found: header_name(header),
}),
}
}
deserialize_primitive!(deserialize_i8, I8, visit_i8, get_i8_value);
deserialize_primitive!(deserialize_i16, I16, visit_i16, get_i16_value);
deserialize_primitive!(deserialize_i32, I32, visit_i32, get_i32_value);
deserialize_primitive!(deserialize_i64, I64, visit_i64, get_i64_value);
deserialize_primitive!(deserialize_i128, I128, visit_i128, get_i128_value);
deserialize_primitive!(deserialize_u8, U8, visit_u8, get_u8_value);
deserialize_primitive!(deserialize_u16, U16, visit_u16, get_u16_value);
deserialize_primitive!(deserialize_u32, U32, visit_u32, get_u32_value);
deserialize_primitive!(deserialize_u64, U64, visit_u64, get_u64_value);
deserialize_primitive!(deserialize_u128, U128, visit_u128, get_u128_value);
deserialize_primitive!(deserialize_f32, F32, visit_f32, get_f32_value);
deserialize_primitive!(deserialize_f64, F64, visit_f64, get_f64_value);
deserialize_primitive!(deserialize_string, STRING, visit_string, get_string_value);
fn deserialize_seq<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let kind = match self.get_byte()? {
STRING_ARRAY => ArrayKind::String,
BOOL_ARRAY => ArrayKind::Boolean,
I8_ARRAY => ArrayKind::I8,
I16_ARRAY => ArrayKind::I16,
I32_ARRAY => ArrayKind::I32,
I64_ARRAY => ArrayKind::I64,
I128_ARRAY => ArrayKind::I128,
U8_ARRAY => ArrayKind::U8,
U16_ARRAY => ArrayKind::U16,
U32_ARRAY => ArrayKind::U32,
U64_ARRAY => ArrayKind::U64,
U128_ARRAY => ArrayKind::U128,
F32_ARRAY => ArrayKind::F32,
F64_ARRAY => ArrayKind::F64,
GENERIC_ARRAY => ArrayKind::Generic,
header => {
return Err(Error::WrongType {
expected: "array",
found: header_name(header),
});
}
};
let size = self.get_size()?;
visitor.visit_seq(SeqDeserializer::new(self, size, kind))
}
fn deserialize_map<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let kind = match self.get_byte()? {
STRING_OBJECT => ObjectKind::String,
I8_OBJECT => ObjectKind::I8,
I16_OBJECT => ObjectKind::I16,
I32_OBJECT => ObjectKind::I32,
I64_OBJECT => ObjectKind::I64,
I128_OBJECT => ObjectKind::I128,
U8_OBJECT => ObjectKind::U8,
U16_OBJECT => ObjectKind::U16,
U32_OBJECT => ObjectKind::U32,
U64_OBJECT => ObjectKind::U64,
U128_OBJECT => ObjectKind::U128,
header => {
return Err(Error::WrongType {
expected: "object",
found: header_name(header),
});
}
};
let size = self.get_size()?;
visitor.visit_map(MapDeserializer::new(self, size, kind))
}
fn deserialize_enum<V>(
self,
_name: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
visitor.visit_enum(EnumDeserializer { deserializer: self })
}
fn deserialize_newtype_struct<V>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
visitor.visit_newtype_struct(self)
}
fn deserialize_bytes<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
visitor.visit_bytes(&self.get_u8_array()?)
}
fn deserialize_byte_buf<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
visitor.visit_byte_buf(self.get_u8_array()?)
}
fn deserialize_char<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self.get_byte()? {
STRING => {}
header => {
return Err(Error::WrongType {
expected: header_name(STRING),
found: header_name(header),
});
}
}
let string = self.get_string_value()?;
if string.len() != 1 {
return Err(Error::WrongType {
expected: "character",
found: header_name(STRING),
});
}
let Some(char) = string.chars().next() else {
return Err(Error::NoChar);
};
visitor.visit_char(char)
}
fn deserialize_str<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self.get_byte()? {
STRING => {}
header => {
return Err(Error::WrongType {
expected: "string",
found: header_name(header),
});
}
}
let string = self.get_string_value()?;
visitor.visit_str(&string)
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self.peek_byte()? {
NULL => {
self.get_byte()?;
visitor.visit_none()
}
_ => visitor.visit_some(self),
}
}
fn deserialize_unit<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self.get_byte()? {
NULL => visitor.visit_unit(),
header => Err(Error::WrongType {
expected: header_name(NULL),
found: header_name(header),
}),
}
}
forward_to_deserialize_any! {
unit_struct tuple tuple_struct struct identifier ignored_any
}
}
pub fn from_reader<T: serde::de::DeserializeOwned>(reader: impl Read) -> Result<T, Error> {
let mut deserializer = Deserializer::new(reader);
T::deserialize(&mut deserializer)
}
pub fn from_bytes<'de, T: serde::de::Deserialize<'de>>(bytes: &'de [u8]) -> Result<T, Error> {
let mut deserializer = Deserializer::new(bytes);
T::deserialize(&mut deserializer)
}