diff --git a/src/de.rs b/src/de.rs index 0387c17..7b00c2f 100644 --- a/src/de.rs +++ b/src/de.rs @@ -1,7 +1,7 @@ use serde; use serde::de::IntoDeserializer; -use rlua::{Value, TablePairs, TableSequence}; +use rlua::{Value, Table, TablePairs, TableSequence}; use error::{Error, Result}; @@ -23,16 +23,17 @@ impl<'lua, 'de> serde::Deserializer<'de> for Deserializer<'lua> { Value::Integer(v) => visitor.visit_i64(v), Value::Number(v) => visitor.visit_f64(v), Value::String(v) => visitor.visit_str(v.to_str()?), - Value::Table(v) => { - let len = v.len()? as usize; - let mut deserializer = MapDeserializer(v.pairs(), None); - let map = visitor.visit_map(&mut deserializer)?; - let remaining = deserializer.0.count(); - if remaining == 0 { - Ok(map) - } else { - Err(serde::de::Error::invalid_length(len, &"fewer elements in array")) + Value::Table(ref v) => { + let mut len = v.clone().pairs::().count(); + for i in 1..=len { + if !v.contains_key(i)? { + len += 1 + } } + if len == 0 || v.len()? as usize != len { + return self.deserialize_map(visitor) + } + self.deserialize_tuple(len, visitor) }, _ => Err(serde::de::Error::custom("invalid value type")), } @@ -101,23 +102,51 @@ impl<'lua, 'de> serde::Deserializer<'de> for Deserializer<'lua> { } #[inline] - fn deserialize_tuple(self, _len: usize, visitor: V) -> Result + fn deserialize_tuple(self, len: usize, visitor: V) -> Result + where V: serde::de::Visitor<'de> + { + match self.value { + Value::Table(v) => { + let mut deserializer = TupleDeserializer(v, 1..=len); + let seq = visitor.visit_seq(&mut deserializer)?; + Ok(seq) + } + _ => Err(serde::de::Error::custom("invalid value type")), + } + } + + #[inline] + fn deserialize_tuple_struct(self, _name: &'static str, len: usize, visitor: V) -> Result where V: serde::de::Visitor<'de> { - self.deserialize_seq(visitor) + self.deserialize_tuple(len, visitor) } #[inline] - fn deserialize_tuple_struct(self, _name: &'static str, _len: usize, visitor: V) -> Result + fn deserialize_map(self, visitor: V) -> Result where V: serde::de::Visitor<'de> { - self.deserialize_seq(visitor) + match self.value { + Value::Table(v) => { + let mut deserializer = MapDeserializer(v.pairs(), None); + let map = visitor.visit_map(&mut deserializer)?; + Ok(map) + } + _ => Err(serde::de::Error::custom("invalid value type")), + } + } + + #[inline] + fn deserialize_struct(self, _name: &'static str, _fields: &'static [&'static str], visitor: V) -> Result + where V: serde::de::Visitor<'de> + { + self.deserialize_map(visitor) } forward_to_deserialize_any! { bool i8 i16 i32 i64 u8 u16 u32 u64 f32 f64 char str string bytes byte_buf unit unit_struct newtype_struct - map struct identifier ignored_any + identifier ignored_any } } @@ -146,6 +175,27 @@ impl<'lua, 'de> serde::de::SeqAccess<'de> for SeqDeserializer<'lua> { } +struct TupleDeserializer<'lua>(Table<'lua>, std::ops::RangeInclusive); + +impl<'lua, 'de> serde::de::SeqAccess<'de> for TupleDeserializer<'lua> { + type Error = Error; + + fn next_element_seed(&mut self, seed: T) -> Result> + where T: serde::de::DeserializeSeed<'de> + { + match self.1.next() { + Some(i) => seed.deserialize(Deserializer { value: self.0.get(i)? }) + .map(Some), + None => Ok(None) + } + } + + fn size_hint(&self) -> Option { + Some(*self.1.end()) + } +} + + struct MapDeserializer<'lua>( TablePairs<'lua, Value<'lua>, Value<'lua>>, Option> @@ -268,7 +318,6 @@ impl<'lua, 'de> serde::de::VariantAccess<'de> for VariantDeserializer<'lua> { #[cfg(test)] mod tests { use rlua::Lua; - use from_value; #[test] @@ -304,6 +353,76 @@ mod tests { }); } + #[test] + fn test_any() { + use serde::private::de::Content; + let lua = Lua::new(); + lua.context(|lua| { + let expected = Content::Map(vec![]); + let value = lua.load( + r#" + a = {} + return a + "#).eval().unwrap(); + let got = from_value::(value).unwrap(); + assert_eq!(format!("{:?}", expected), format!("{:?}", got)); + + let expected = Content::Seq(vec![Content::I64(10), Content::I64(20), Content::I64(30), Content::I64(40)]); + let value = lua.load( + r#" + a = {10, 20, 30} + a[4] = 40 + return a + "#).eval().unwrap(); + let got = from_value::(value).unwrap(); + assert_eq!(format!("{:?}", expected), format!("{:?}", got)); + + let expected = Content::Seq(vec![Content::I64(10), Content::Unit, Content::I64(30)]); + let value = lua.load( + r#" + a = {10, 20, 30} + a[2] = nil + return a + "#).eval().unwrap(); + let got = from_value::(value).unwrap(); + assert_eq!(format!("{:?}", expected), format!("{:?}", got)); + + let expected = Content::Seq(vec![Content::I64(10), Content::I64(20)]); + let value = lua.load( + r#" + a = {10, 20, 30} + a[3] = nil + return a + "#).eval().unwrap(); + let got = from_value::(value).unwrap(); + assert_eq!(format!("{:?}", expected), format!("{:?}", got)); + + let expected = Content::Map(vec![(Content::I64(1), Content::I64(10)), (Content::I64(3), Content::I64(30))]); + let value = lua.load( + r#" + a = {} + a[1] = 10 + a[3] = 30 + return a + "#).eval().unwrap(); + let got = from_value::(value).unwrap(); + assert_eq!(format!("{:?}", expected), format!("{:?}", got)); + + let expected = Content::Map(vec![ + (Content::I64(2), Content::I64(20)), + (Content::I64(3), Content::I64(30)), + (Content::String("a".to_string()), Content::I64(1)), + ]); + let value = lua.load( + r#" + a = {nil, 20, 30, a=1} + return a + "#).eval().unwrap(); + let got = from_value::(value).unwrap(); + assert_eq!(format!("{:?}", expected), format!("{:?}", got)); + }); + } + #[test] fn test_tuple() { #[derive(Deserialize, PartialEq, Debug)] @@ -337,6 +456,7 @@ mod tests { enum E { Unit, Newtype(u32), + Array(Vec), Tuple(u32, u32), Struct { a: u32 }, } @@ -351,7 +471,6 @@ mod tests { let got = from_value(value).unwrap(); assert_eq!(expected, got); - let expected = E::Newtype(1); let value = lua.load( r#" @@ -362,6 +481,16 @@ mod tests { let got = from_value(value).unwrap(); assert_eq!(expected, got); + let expected = E::Array(vec![10, 20, 30]); + let value = lua.load( + r#" + a = {} + a["Array"] = {10, 20, 30} + return a + "#).eval().unwrap(); + let got = from_value(value).unwrap(); + assert_eq!(expected, got); + let expected = E::Tuple(1, 2); let value = lua.load( r#" @@ -384,4 +513,73 @@ mod tests { assert_eq!(expected, got); }); } + + #[test] + fn test_enum_untagged() { + #[derive(Deserialize, PartialEq, Debug)] + #[serde(untagged)] + enum E { + Unit, + Newtype(u32), + Tuple(u32, u32), + Array(Vec), + Struct { a: u32 }, + Table(std::collections::HashMap), + } + + let lua = Lua::new(); + lua.context(|lua| { + let expected = E::Unit; + let value = lua.load( + r#" + return + "#).eval().unwrap(); + let got = from_value(value).unwrap(); + assert_eq!(expected, got); + + let expected = E::Newtype(1); + let value = lua.load( + r#" + return 1 + "#).eval().unwrap(); + let got = from_value(value).unwrap(); + assert_eq!(expected, got); + + let expected = E::Tuple(1, 2); + let value = lua.load( + r#" + return {1, 2} + "#).eval().unwrap(); + let got = from_value(value).unwrap(); + assert_eq!(expected, got); + + let expected = E::Array(vec![10, 20, 30]); + let value = lua.load( + r#" + return {10, 20, 30} + "#).eval().unwrap(); + let got = from_value(value).unwrap(); + assert_eq!(expected, got); + + let expected = E::Struct { a: 1 }; + let value = lua.load( + r#" + a = {} + a["a"] = 1 + return a + "#).eval().unwrap(); + let got = from_value(value).unwrap(); + assert_eq!(expected, got); + + let expected = E::Table(vec![("b".to_string(), 3)].into_iter().collect()); + let value = lua.load( + r#" + a = {} + a["b"] = 3 + return a + "#).eval().unwrap(); + let got = from_value(value).unwrap(); + assert_eq!(expected, got); + }); + } } diff --git a/src/ser.rs b/src/ser.rs index a12e837..f7a8941 100644 --- a/src/ser.rs +++ b/src/ser.rs @@ -374,9 +374,16 @@ mod tests { struct Test { int: u32, seq: Vec<&'static str>, + map: std::collections::HashMap, + empty: Vec<()>, } - let test = Test { int: 1, seq: vec!["a", "b"] }; + let test = Test { + int: 1, + seq: vec!["a", "b"], + map: vec![(1, 2), (4, 1)].into_iter().collect(), + empty: vec![] + }; let lua = Lua::new(); lua.context(|lua| { @@ -387,16 +394,20 @@ mod tests { assert(value["int"] == 1) assert(value["seq"][1] == "a") assert(value["seq"][2] == "b") + assert(value["map"][1] == 2) + assert(value["map"][4] == 1) + assert(next(value["empty"]) == nil) "#).exec() }).unwrap() } #[test] - fn test_num() { + fn test_enum() { #[derive(Serialize)] enum E { Unit, Newtype(u32), + Array(Vec), Tuple(u32, u32), Struct { a: u32}, } @@ -418,6 +429,15 @@ mod tests { assert(value["Newtype"] == 1) "#).exec().unwrap(); + let t = E::Array(vec![10, 20, 30]); + let value = to_value(lua, &t).unwrap(); + lua.globals().set("value", value).unwrap(); + lua.load(r#" + assert(value["Array"][1] == 10) + assert(value["Array"][2] == 20) + assert(value["Array"][3] == 30) + "#).exec().unwrap(); + let t = E::Tuple(1, 2); let value = to_value(lua, &t).unwrap(); lua.globals().set("value", value).unwrap(); @@ -434,4 +454,66 @@ mod tests { "#).exec() }).unwrap(); } + + #[test] + fn test_enum_untagged() { + #[derive(Serialize)] + #[serde(untagged)] + enum E { + Unit, + Newtype(u32), + Tuple(u32, u32), + Array(Vec), + Struct { a: u32 }, + Table(std::collections::HashMap), + } + + let lua = Lua::new(); + lua.context(|lua| { + let u = E::Unit; + let value = to_value(lua, &u).unwrap(); + lua.globals().set("value", value).unwrap(); + lua.load(r#" + assert(value == nil) + "#).exec().unwrap(); + + let n = E::Newtype(1); + let value = to_value(lua, &n).unwrap(); + lua.globals().set("value", value).unwrap(); + lua.load(r#" + assert(value == 1) + "#).exec().unwrap(); + + let t = E::Tuple(1, 2); + let value = to_value(lua, &t).unwrap(); + lua.globals().set("value", value).unwrap(); + lua.load(r#" + assert(value[1] == 1) + assert(value[2] == 2) + "#).exec().unwrap(); + + let t = E::Array(vec![10, 20, 30]); + let value = to_value(lua, &t).unwrap(); + lua.globals().set("value", value).unwrap(); + lua.load(r#" + assert(value[1] == 10) + assert(value[2] == 20) + assert(value[3] == 30) + "#).exec().unwrap(); + + let s = E::Struct { a: 1 }; + let value = to_value(lua, &s).unwrap(); + lua.globals().set("value", value).unwrap(); + lua.load(r#" + assert(value["a"] == 1) + "#).exec().unwrap(); + + let s = E::Table(vec![("b".to_string(), 3)].into_iter().collect()); + let value = to_value(lua, &s).unwrap(); + lua.globals().set("value", value).unwrap(); + lua.load(r#" + assert(value["b"] == 3) + "#).exec() + }).unwrap(); + } }