diff --git a/src/de/map.rs b/src/de/map.rs new file mode 100644 index 0000000..837fc31 --- /dev/null +++ b/src/de/map.rs @@ -0,0 +1,54 @@ +use serde::de::{DeserializeSeed, MapAccess}; + +use de::{Deserializer, Result}; +use error::Error; +use node_types::StandardType; + +pub struct Map<'a, 'de: 'a> { + de: &'a mut Deserializer<'de>, +} + +impl<'de, 'a> Map<'a, 'de> { + pub fn new(de: &'a mut Deserializer<'de>) -> Self { + Self { de } + } +} + +impl<'de, 'a> MapAccess<'de> for Map<'a, 'de> { + type Error = Error; + + fn next_key_seed(&mut self, seed: K) -> Result> + where K: DeserializeSeed<'de> + { + trace!("--> ::next_key_seed()"); + + let (node_type, _is_array) = self.de.reader.read_node_type()?; + debug!("::next_key_seed() => node_type: {:?}", node_type); + + if node_type == StandardType::NodeEnd { + trace!("::next_key_seed() => end of map"); + return Ok(None); + } + + let value = seed.deserialize(&mut *self.de).map(Some)?; + + /* + if node_type != StandardType::NodeStart { + // Consume the end node and do a sanity check + let node_type = self.de.read_node()?; + if node_type != StandardType::NodeEnd { + return Err(KbinErrorKind::TypeMismatch(*StandardType::NodeEnd, *node_type).into()); + } + } + */ + + Ok(value) + } + + fn next_value_seed(&mut self, seed: V) -> Result + where V: DeserializeSeed<'de> + { + debug!("--> ::next_value_seed()"); + seed.deserialize(&mut *self.de) + } +} diff --git a/src/de/mod.rs b/src/de/mod.rs index 2c4f03d..dadad50 100644 --- a/src/de/mod.rs +++ b/src/de/mod.rs @@ -4,17 +4,15 @@ use byteorder::{BigEndian, ByteOrder, ReadBytesExt}; use failure::ResultExt; use serde::de::{self, Deserialize, Visitor}; -use byte_buffer::ByteBufferRead; -use compression::Compression; -use encoding_type::EncodingType; use error::{Error, KbinErrorKind}; use node_types::StandardType; -use sixbit::unpack_sixbit; -use super::{ARRAY_MASK, SIGNATURE, SIG_COMPRESSED}; +use reader::Reader; +mod map; mod seq; mod structure; +use self::map::Map; use self::seq::Seq; use self::structure::Struct; @@ -26,14 +24,11 @@ enum ReadMode { } pub struct Deserializer<'de> { - encoding: EncodingType, - read_mode: ReadMode, + node_stack: Vec, first_struct: bool, - //node_buf_end: u64, - node_buf: ByteBufferRead<&'de [u8]>, - data_buf: ByteBufferRead<&'de [u8]>, + reader: Reader<'de>, } pub fn from_bytes<'a, T>(input: &'a [u8]) -> Result @@ -46,76 +41,22 @@ pub fn from_bytes<'a, T>(input: &'a [u8]) -> Result impl<'de> Deserializer<'de> { pub fn new(input: &'de [u8]) -> Result { - // Node buffer starts from the beginning. - // Data buffer starts later after reading `len_data`. - let mut node_buf = ByteBufferRead::new(&input[..]); - - let signature = node_buf.read_u8().context(KbinErrorKind::HeaderRead("signature"))?; - if signature != SIGNATURE { - return Err(KbinErrorKind::HeaderValue("signature").into()); - } - - // TODO: support uncompressed - let compress_byte = node_buf.read_u8().context(KbinErrorKind::HeaderRead("compression"))?; - if compress_byte != SIG_COMPRESSED { - return Err(KbinErrorKind::HeaderValue("compression").into()); - } - - let compressed = Compression::from_byte(compress_byte)?; - - let encoding_byte = node_buf.read_u8().context(KbinErrorKind::HeaderRead("encoding"))?; - let encoding_negation = node_buf.read_u8().context(KbinErrorKind::HeaderRead("encoding negation"))?; - let encoding = EncodingType::from_byte(encoding_byte)?; - if encoding_negation != !encoding_byte { - return Err(KbinErrorKind::HeaderValue("encoding negation").into()); - } - - info!("signature: 0x{:x}, compression: 0x{:x} ({:?}), encoding: 0x{:x} ({:?})", signature, compress_byte, compressed, encoding_byte, encoding); - - let len_node = node_buf.read_u32::().context(KbinErrorKind::LenNodeRead)?; - info!("len_node: {0} (0x{0:x})", len_node); - - // We have read 8 bytes so far, so offset the start of the data buffer from - // the start of the input data. - let data_buf_start = len_node + 8; - let mut data_buf = ByteBufferRead::new(&input[(data_buf_start as usize)..]); - - let len_data = data_buf.read_u32::().context(KbinErrorKind::LenDataRead)?; - info!("len_data: {0} (0x{0:x})", len_data); - - //let node_buf_end = data_buf_start.into(); + let reader = Reader::new(input)?; Ok(Self { - encoding, read_mode: ReadMode::Single, first_struct: true, - //node_buf_end, - node_buf, - data_buf, + node_stack: Vec::new(), + reader, }) } - fn read_node(&mut self) -> Result { - let raw_node_type = self.node_buf.read_u8().context(KbinErrorKind::NodeTypeRead)?; - let is_array = raw_node_type & ARRAY_MASK == ARRAY_MASK; - let node_type = raw_node_type & !ARRAY_MASK; - - let xml_type = StandardType::from_u8(node_type); - debug!("raw_node_type: {}, node_type: {:?} ({}), is_array: {}", raw_node_type, xml_type, node_type, is_array); - - Ok(xml_type) - } - - fn read_name(&mut self) -> Result { - unpack_sixbit(&mut *self.node_buf).map_err(Error::from) - } - - fn read_node_with_name(&mut self) -> Result<(StandardType, String)> { - let node_type = self.read_node()?; - let name = self.read_name()?; + fn read_node_with_name(&mut self) -> Result<(StandardType, bool, String)> { + let (node_type, is_array) = self.reader.read_node_type()?; + let name = self.reader.read_node_identifier()?; debug!("name: {}", name); - Ok((node_type, name)) + Ok((node_type, is_array, name)) } } @@ -126,10 +67,10 @@ macro_rules! de_type { { let value = match self.read_mode { ReadMode::Single => { - self.data_buf.get_aligned(*StandardType::$standard_type)?[0] $($cast)* + self.reader.data_buf.get_aligned(*StandardType::$standard_type)?[0] $($cast)* }, ReadMode::Array => { - self.data_buf.read_u8().context(KbinErrorKind::DataRead(1))? $($cast)* + self.reader.read_u8().context(KbinErrorKind::DataRead(1))? $($cast)* }, }; trace!(concat!("Deserializer::", stringify!($method), "() => value: {:?}"), value); @@ -143,11 +84,11 @@ macro_rules! de_type { { let value = match self.read_mode { ReadMode::Single => { - let value = self.data_buf.get_aligned(*StandardType::$standard_type)?; + let value = self.reader.data_buf.get_aligned(*StandardType::$standard_type)?; BigEndian::$read_method(&value) }, ReadMode::Array => { - self.data_buf.$read_method::().context(KbinErrorKind::DataRead(StandardType::$standard_type.size as usize))? + self.reader.data_buf.$read_method::().context(KbinErrorKind::DataRead(StandardType::$standard_type.size as usize))? }, }; trace!(concat!("Deserializer::", stringify!($method), "() => value: {:?}"), value); @@ -175,11 +116,31 @@ impl<'de, 'a> de::Deserializer<'de> for &'a mut Deserializer<'de> { false } - fn deserialize_any(self, _visitor: V) -> Result + fn deserialize_any(self, visitor: V) -> Result where V: Visitor<'de> { trace!("Deserializer::deserialize_any()"); - Err(KbinErrorKind::DataRead(1).into()) + + let (node_type, _is_array) = self.reader.peek_node_type()?; + debug!("Deserializer::deserialize_any() => node_type: {:?}", node_type); + + let value = match node_type { + StandardType::Attribute | + StandardType::NodeStart => self.deserialize_identifier(visitor), + StandardType::Binary => self.deserialize_bytes(visitor), + StandardType::String => self.deserialize_string(visitor), + StandardType::U8 => self.deserialize_u8(visitor), + StandardType::U16 => self.deserialize_u16(visitor), + StandardType::U32 => self.deserialize_u32(visitor), + StandardType::U64 => self.deserialize_u64(visitor), + StandardType::S8 => self.deserialize_i8(visitor), + StandardType::S16 => self.deserialize_i16(visitor), + StandardType::S32 => self.deserialize_i32(visitor), + StandardType::S64 => self.deserialize_i64(visitor), + StandardType::NodeEnd => visitor.visit_none(), + _ => unimplemented!(), + }; + value } fn deserialize_bool(self, visitor: V) -> Result @@ -187,7 +148,7 @@ impl<'de, 'a> de::Deserializer<'de> for &'a mut Deserializer<'de> { { trace!("Deserializer::deserialize_bool()"); - let value = self.data_buf.get_aligned(*StandardType::Boolean)?[0]; + let value = self.reader.data_buf.get_aligned(*StandardType::Boolean)?[0]; trace!("Deserializer::deserialize_bool() => value: {:?}", value); let value = match value { @@ -217,7 +178,7 @@ impl<'de, 'a> de::Deserializer<'de> for &'a mut Deserializer<'de> { { trace!("Deserializer::deserialize_string()"); - visitor.visit_string(self.data_buf.read_str(self.encoding)?) + visitor.visit_string(self.reader.read_string()?) } implement_type!(deserialize_bytes); @@ -253,16 +214,26 @@ impl<'de, 'a> de::Deserializer<'de> for &'a mut Deserializer<'de> { { trace!("Deserializer::deserialize_seq()"); - // TODO: add size check against len - let size = self.data_buf.read_u32::().context(KbinErrorKind::ArrayLengthRead)?; - debug!("Deserializer::deserialize_seq() => read array size: {}", size); + let node_type = self.node_stack.last().ok_or(KbinErrorKind::InvalidState)?.clone(); - // Changes to `self.read_mode` must stay here as `next_element_seed` is not - // called past the length of the array to reset the read mode - self.read_mode = ReadMode::Array; - let value = visitor.visit_seq(Seq::new(self, size as usize))?; - self.read_mode = ReadMode::Single; - self.data_buf.realign_reads(None)?; + // If the last node type on the stack is a `NodeStart` then we are likely + // collecting a list of structs + let value = if node_type == StandardType::NodeStart { + visitor.visit_seq(Seq::new(self, None))? + } else { + // TODO: add size check against len + let size = self.reader.read_u32().context(KbinErrorKind::ArrayLengthRead)?; + debug!("Deserializer::deserialize_seq() => read array size: {}", size); + + // Changes to `self.read_mode` must stay here as `next_element_seed` is not + // called past the length of the array to reset the read mode + self.read_mode = ReadMode::Array; + let value = visitor.visit_seq(Seq::new(self, Some(size as usize)))?; + self.read_mode = ReadMode::Single; + self.reader.data_buf.realign_reads(None)?; + + value + }; Ok(value) } @@ -281,18 +252,22 @@ impl<'de, 'a> de::Deserializer<'de> for &'a mut Deserializer<'de> { trace!("Deserializer::deserialize_tuple_struct(name: {:?}, len: {})", name, len); self.read_mode = ReadMode::Array; - let value = visitor.visit_seq(Seq::new(self, len))?; + let value = visitor.visit_seq(Seq::new(self, Some(len)))?; self.read_mode = ReadMode::Single; - self.data_buf.realign_reads(None)?; + self.reader.data_buf.realign_reads(None)?; Ok(value) } - fn deserialize_map(self, _visitor: V) -> Result + fn deserialize_map(self, visitor: V) -> Result where V: Visitor<'de> { trace!("Deserializer::deserialize_map()"); - unimplemented!(); + + let (node_type, _, name) = self.read_node_with_name()?; + debug!("Deserializer::deserialize_map() => node_type: {:?}, name: {:?}", node_type, name); + + visitor.visit_map(Map::new(self)) } fn deserialize_struct(self, name: &'static str, fields: &'static [&'static str], visitor: V) -> Result @@ -304,8 +279,8 @@ impl<'de, 'a> de::Deserializer<'de> for &'a mut Deserializer<'de> { // The `NodeStart` event is consumed by `deserialize_identifier` when // reading the parent struct, don't consume the next event. if self.first_struct { - let (node_type, name) = self.read_node_with_name()?; - debug!("node_type: {:?}, name: {:?}", node_type, name); + let (node_type, _, name) = self.read_node_with_name()?; + debug!("Deserializer::deserialize_struct() => node_type: {:?}, name: {:?}", node_type, name); // Sanity check if node_type != StandardType::NodeStart { @@ -331,10 +306,19 @@ impl<'de, 'a> de::Deserializer<'de> for &'a mut Deserializer<'de> { { trace!("Deserializer::deserialize_identifier()"); + let name = self.reader.read_node_identifier()?; + debug!("Deserializer::deserialize_identifier() => name: {}", name); + // Do not use `deserialize_string`! That reads from the data buffer and // this reads a sixbit string from the node buffer - visitor.visit_string(self.read_name()?) + visitor.visit_string(name) } - implement_type!(deserialize_ignored_any); + fn deserialize_ignored_any(self, visitor: V) -> Result + where V: Visitor<'de> + { + trace!("Deserializer::deserialize_ignored_any()"); + + self.deserialize_any(visitor) + } } diff --git a/src/de/structure.rs b/src/de/structure.rs index d66ce5c..e313495 100644 --- a/src/de/structure.rs +++ b/src/de/structure.rs @@ -24,33 +24,50 @@ impl<'de, 'a> MapAccess<'de> for Struct<'a, 'de> { fn next_key_seed(&mut self, seed: K) -> Result> where K: DeserializeSeed<'de> { - trace!("MapAccess::next_key_seed()"); + trace!("--> ::next_key_seed()"); - let node_type = self.de.read_node()?; - debug!("MapAccess::next_key_seed() => node_type: {:?}", node_type); + let (node_type, _is_array) = self.de.reader.read_node_type()?; + debug!("Struct::next_key_seed() => node_type: {:?}", node_type); if node_type == StandardType::NodeEnd { - trace!("MapAccess::next_key_seed() => end of map"); + trace!("Struct::next_key_seed() => end of map"); return Ok(None); } let value = seed.deserialize(&mut *self.de).map(Some)?; - if node_type != StandardType::NodeStart { - // Consume the end node and do a sanity check - let node_type = self.de.read_node()?; - if node_type != StandardType::NodeEnd { - return Err(KbinErrorKind::TypeMismatch(*StandardType::NodeEnd, *node_type).into()); - } + match node_type { + StandardType::NodeStart => { + debug!("Struct::next_key_seed() => got a node start!"); + }, + StandardType::Attribute => { + debug!("Struct::next_key_seed() => got an attribute!"); + }, + _ => { + // Consume the end node and do a sanity check + let (node_type, _is_array) = self.de.reader.read_node_type()?; + if node_type != StandardType::NodeEnd { + return Err(KbinErrorKind::TypeMismatch(*StandardType::NodeEnd, *node_type).into()); + } + }, } + // Store the current node type on the stack for stateful handling based on + // the current node type + self.de.node_stack.push(node_type); + Ok(value) } fn next_value_seed(&mut self, seed: V) -> Result where V: DeserializeSeed<'de> { - debug!("MapAccess::next_value_seed()"); - seed.deserialize(&mut *self.de) + debug!("--> ::next_value_seed()"); + let value = seed.deserialize(&mut *self.de)?; + + let popped = self.de.node_stack.pop(); + debug!("::next_value_seed() => popped: {:?}, node_stack: {:?}", popped, self.de.node_stack); + + Ok(value) } } diff --git a/src/error.rs b/src/error.rs index 0f3c0f9..10dd799 100644 --- a/src/error.rs +++ b/src/error.rs @@ -103,6 +103,9 @@ pub enum KbinErrorKind { #[fail(display = "Type mismatch, expected: {}, found: {}", _0, _1)] TypeMismatch(KbinType, KbinType), + + #[fail(display = "Invalid state")] + InvalidState, } impl fmt::Display for KbinError {