//! Decoding: reading protocol values from a byte stream. use std::{ collections::{BTreeMap, BTreeSet, BinaryHeap, HashMap, HashSet}, hash::Hash, io::Read, sync::Arc, }; use crate::{DEFAULT_BUFFER_LEN, codec::error::Result}; use crate::{codec::error::CodecError, context::Context}; macro_rules! impl_decode { (u8) => { impl Decode for u8 { fn decode(reader: &mut dyn Read, _: &Context) -> Result where Self: Sized, { let mut byte = [0u8; 1]; reader.read_exact(&mut byte)?; Ok(byte[0]) } fn decode_slice(reader: &mut dyn Read, _: &Context) -> Result> where Self: Sized, { let mut buf = Vec::new(); reader.read_to_end(&mut buf)?; Ok(buf) } } }; ($type: ty) => { impl Decode for $type { // decode number using big endian fn decode(reader: &mut dyn Read, _: &Context) -> Result where Self: Sized, { const BYTES: usize = (<$type>::BITS / 8) as usize; let mut bytes = [0; BYTES]; reader.read_exact(&mut bytes)?; Ok(<$type>::from_be_bytes(bytes)) } fn decode_slice(reader: &mut dyn Read, _: &Context) -> Result> where Self: Sized, { let mut buf = Vec::with_capacity(DEFAULT_BUFFER_LEN); reader.read_to_end(&mut buf); const BYTES: usize = <$type>::BITS as usize / 8; Ok(buf .chunks(BYTES) .map(|slice| { let Some(bytes): Option<&[u8; BYTES]> = slice.as_array() else { unreachable!() }; <$type>::from_be_bytes(*bytes) }) .collect()) } } }; } macro_rules! impl_decode_tuples { ($($generic: ident),+) => { impl Decode for ($($generic,)+) where $($generic: Decode,)+ { #[allow(non_snake_case)] fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized { $( let $generic = $generic::decode(reader, ctx)?; )+ Ok(($($generic,)+)) } } }; } /// Read a value from a byte stream pub trait Decode { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized; // TODO : change for anything else than Vec /// assume reader contains only the slice to decode fn decode_slice(reader: &mut dyn Read, ctx: &Context) -> Result> where Self: Sized, { let mut buf = Vec::new(); loop { match Self::decode(reader, ctx) { Ok(value) => buf.push(value), Err(CodecError::IoError(err)) if let std::io::ErrorKind::UnexpectedEof = err.kind() => { return Ok(buf); } Err(err) => return Err(err), } } } } // ~ Decode arrays impl> Decode for Vec { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let len = usize::decode(reader, ctx)?; let mut vec = Vec::with_capacity(len); for _ in 0..len { vec.push(T::decode(reader, ctx)?); } Ok(vec) } } impl> Decode for Option { /// Assume is always some fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { T::decode(reader, ctx).map(Some) } } // ~ Decode slices impl Decode for [T; S] where T: Decode + Default + Copy, { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let slice = T::decode_slice(reader, ctx)?; let got = slice.len(); slice .as_array() .cloned() .ok_or(CodecError::WrongSize { expected: S, got }) } } impl Decode for &[T] where T: Decode + Default + Copy, { fn decode(_: &mut dyn Read, _: &Context) -> Result where Self: Sized, { unimplemented!() } } // ~ Decode set impl Decode for HashSet where T: Decode + Eq + Hash, { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let slice = T::decode_slice(reader, ctx)?; Ok(slice.into_iter().collect()) } } impl Decode for BTreeSet where T: Decode + Ord, { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let slice = T::decode_slice(reader, ctx)?; Ok(slice.into_iter().collect()) } } // ~ Decode misc impl Decode for BinaryHeap where T: Decode + Ord, { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let slice = T::decode_slice(reader, ctx)?; Ok(BinaryHeap::from_iter(slice)) } } // ~ Decode maps impl Decode for HashMap where K: Decode + Eq + Hash, V: Decode, { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let slice = <(K, V) as Decode>::decode_slice(reader, ctx)?; Ok(slice.into_iter().collect()) } } impl Decode for BTreeMap where K: Decode + Ord, V: Decode, { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let slice = <(K, V) as Decode>::decode_slice(reader, ctx)?; Ok(slice.into_iter().collect()) } } // ~ Decode string impl Decode for &str { fn decode(_: &mut dyn Read, _: &Context) -> Result where Self: Sized, { unimplemented!() } } impl Decode for String { fn decode(reader: &mut dyn Read, _: &Context) -> Result where Self: Sized, { let mut bytes: Vec = Vec::new(); reader.read_to_end(&mut bytes)?; Ok(String::from_utf8_lossy(&bytes).to_string()) } } // ~ Decode primitive impl_decode!(u8); impl_decode!(u16); impl_decode!(u32); impl_decode!(u64); impl_decode!(u128); impl_decode!(usize); impl_decode!(i8); impl_decode!(i16); impl_decode!(i32); impl_decode!(i64); impl_decode!(i128); impl_decode!(isize); impl Decode for f64 { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let bits = u64::decode(reader, ctx)?; let value = f64::from_bits(bits); Ok(value) } } impl Decode for f32 { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let bits = u32::decode(reader, ctx)?; let value = f32::from_bits(bits); Ok(value) } } impl Decode for bool { fn decode(reader: &mut dyn Read, _: &Context) -> Result where Self: Sized, { let mut byte = [0_u8]; reader.read_exact(&mut byte)?; Ok(byte[0] != 0) } } // ~ Decode tuple impl_decode_tuples!(A); impl_decode_tuples!(A, B); impl_decode_tuples!(A, B, C); impl_decode_tuples!(A, B, C, D); impl_decode_tuples!(A, B, C, D, E); impl_decode_tuples!(A, B, C, D, E, F); impl_decode_tuples!(A, B, C, D, E, F, G); impl_decode_tuples!(A, B, C, D, E, F, G, H); impl_decode_tuples!(A, B, C, D, E, F, G, H, I); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O, P); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O, P, Q); // ~ Decode pointer impl Decode for std::sync::Arc where T: Decode, { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result { T::decode(reader, ctx).map(Arc::new) } } impl Decode for Arc<[T]> where T: Decode, { fn decode(reader: &mut dyn Read, ctx: &Context) -> Result where Self: Sized, { let slice = T::decode_slice(reader, ctx)?; Ok(slice.into()) } }