From 3a9c8c2da2ed3d35cc3550417ec8d475a34a3d4f Mon Sep 17 00:00:00 2001 From: Alex Orlenko Date: Tue, 22 Mar 2022 20:52:10 +0000 Subject: [PATCH] Add Luau vector datatype support --- src/conversion.rs | 65 +++++++++++++++++++++++++++-------------------- src/lua.rs | 14 ++++++++++ src/serde/de.rs | 45 ++++++++++++++++++++++++++++++++ src/util.rs | 7 +++++ src/value.rs | 9 +++++++ tests/luau.rs | 20 ++++++++++++++- tests/serde.rs | 23 +++++++++++++++++ 7 files changed, 154 insertions(+), 29 deletions(-) diff --git a/src/conversion.rs b/src/conversion.rs index 67726db..b9939c1 100644 --- a/src/conversion.rs +++ b/src/conversion.rs @@ -464,21 +464,33 @@ impl<'lua, T, const N: usize> FromLua<'lua> for [T; N] where T: FromLua<'lua>, { - fn from_lua(value: Value<'lua>, _: &'lua Lua) -> Result { - if let Value::Table(table) = value { - let vec = table.sequence_values().collect::>>()?; - vec.try_into() - .map_err(|vec: Vec| Error::FromLuaConversionError { - from: "Table", - to: "Array", - message: Some(format!("expected table of length {}, got {}", N, vec.len())), - }) - } else { - Err(Error::FromLuaConversionError { + fn from_lua(value: Value<'lua>, _lua: &'lua Lua) -> Result { + match value { + #[cfg(feature = "luau")] + Value::Vector(x, y, z) if N == 3 => Ok(mlua_expect!( + vec![ + T::from_lua(Value::Number(x as _), _lua)?, + T::from_lua(Value::Number(y as _), _lua)?, + T::from_lua(Value::Number(z as _), _lua)?, + ] + .try_into() + .map_err(|_| ()), + "cannot convert vector to array" + )), + Value::Table(table) => { + let vec = table.sequence_values().collect::>>()?; + vec.try_into() + .map_err(|vec: Vec| Error::FromLuaConversionError { + from: "Table", + to: "Array", + message: Some(format!("expected table of length {}, got {}", N, vec.len())), + }) + } + _ => Err(Error::FromLuaConversionError { from: value.type_name(), to: "Array", message: Some("expected table".to_string()), - }) + }), } } } @@ -490,16 +502,8 @@ impl<'lua, T: ToLua<'lua>> ToLua<'lua> for Box<[T]> { } impl<'lua, T: FromLua<'lua>> FromLua<'lua> for Box<[T]> { - fn from_lua(value: Value<'lua>, _: &'lua Lua) -> Result { - if let Value::Table(table) = value { - table.sequence_values().collect() - } else { - Err(Error::FromLuaConversionError { - from: value.type_name(), - to: "Box<[T]>", - message: Some("expected table".to_string()), - }) - } + fn from_lua(value: Value<'lua>, lua: &'lua Lua) -> Result { + Ok(Vec::::from_lua(value, lua)?.into_boxed_slice()) } } @@ -510,15 +514,20 @@ impl<'lua, T: ToLua<'lua>> ToLua<'lua> for Vec { } impl<'lua, T: FromLua<'lua>> FromLua<'lua> for Vec { - fn from_lua(value: Value<'lua>, _: &'lua Lua) -> Result { - if let Value::Table(table) = value { - table.sequence_values().collect() - } else { - Err(Error::FromLuaConversionError { + fn from_lua(value: Value<'lua>, _lua: &'lua Lua) -> Result { + match value { + #[cfg(feature = "luau")] + Value::Vector(x, y, z) => Ok(vec![ + T::from_lua(Value::Number(x as _), _lua)?, + T::from_lua(Value::Number(y as _), _lua)?, + T::from_lua(Value::Number(z as _), _lua)?, + ]), + Value::Table(table) => table.sequence_values().collect(), + _ => Err(Error::FromLuaConversionError { from: value.type_name(), to: "Vec", message: Some("expected table".to_string()), - }) + }), } } } diff --git a/src/lua.rs b/src/lua.rs index 920801f..2d5625f 100644 --- a/src/lua.rs +++ b/src/lua.rs @@ -1972,6 +1972,11 @@ impl Lua { ffi::lua_pushnumber(self.state, n); } + #[cfg(feature = "luau")] + Value::Vector(x, y, z) => { + ffi::lua_pushvector(self.state, x, y, z); + } + Value::String(s) => { self.push_ref(&s.0); } @@ -2033,6 +2038,15 @@ impl Lua { } } + #[cfg(feature = "luau")] + ffi::LUA_TVECTOR => { + let v = ffi::lua_tovector(state, -1); + mlua_debug_assert!(!v.is_null(), "vector is null"); + let vec = Value::Vector(*v, *v.add(1), *v.add(2)); + ffi::lua_pop(state, 1); + vec + } + ffi::LUA_TSTRING => Value::String(String(self.pop_ref())), ffi::LUA_TTABLE => Value::Table(Table(self.pop_ref())), diff --git a/src/serde/de.rs b/src/serde/de.rs index be6e114..3e4bb14 100644 --- a/src/serde/de.rs +++ b/src/serde/de.rs @@ -123,6 +123,8 @@ impl<'lua, 'de> serde::Deserializer<'de> for Deserializer<'lua> { } #[allow(clippy::useless_conversion)] Value::Number(n) => visitor.visit_f64(n.into()), + #[cfg(feature = "luau")] + Value::Vector(_, _, _) => self.deserialize_seq(visitor), Value::String(s) => match s.to_str() { Ok(s) => visitor.visit_str(s), Err(_) => visitor.visit_bytes(s.as_bytes()), @@ -214,6 +216,16 @@ impl<'lua, 'de> serde::Deserializer<'de> for Deserializer<'lua> { V: de::Visitor<'de>, { match self.value { + #[cfg(feature = "luau")] + Value::Vector(x, y, z) => { + let mut deserializer = VecDeserializer { + vec: [x, y, z], + next: 0, + options: self.options, + visited: self.visited, + }; + visitor.visit_seq(&mut deserializer) + } Value::Table(t) => { let _guard = RecursionGuard::new(&t, &self.visited); @@ -352,6 +364,39 @@ impl<'lua, 'de> de::SeqAccess<'de> for SeqDeserializer<'lua> { } } +#[cfg(feature = "luau")] +struct VecDeserializer { + vec: [f32; 3], + next: usize, + options: Options, + visited: Rc>>, +} + +#[cfg(feature = "luau")] +impl<'de> de::SeqAccess<'de> for VecDeserializer { + type Error = Error; + + fn next_element_seed(&mut self, seed: T) -> Result> + where + T: de::DeserializeSeed<'de>, + { + match self.vec.get(self.next) { + Some(&n) => { + self.next += 1; + let visited = Rc::clone(&self.visited); + let deserializer = + Deserializer::from_parts(Value::Number(n as _), self.options, visited); + seed.deserialize(deserializer).map(Some) + } + None => Ok(None), + } + } + + fn size_hint(&self) -> Option { + Some(3) + } +} + struct MapDeserializer<'lua> { pairs: TablePairs<'lua, Value<'lua>, Value<'lua>>, value: Option>, diff --git a/src/util.rs b/src/util.rs index 0632aa6..66a60db 100644 --- a/src/util.rs +++ b/src/util.rs @@ -965,6 +965,13 @@ pub(crate) unsafe fn to_string(state: *mut ffi::lua_State, index: c_int) -> Stri i.to_string() } } + #[cfg(feature = "luau")] + ffi::LUA_TVECTOR => { + let v = ffi::lua_tovector(state, index); + mlua_debug_assert!(!v.is_null(), "vector is null"); + let (x, y, z) = (*v, *v.add(1), *v.add(2)); + format!("vector({},{},{})", x, y, z) + } ffi::LUA_TSTRING => { let mut size = 0; // This will not trigger a 'm' error, because the reference is guaranteed to be of diff --git a/src/value.rs b/src/value.rs index 290cd4e..1dfcbbd 100644 --- a/src/value.rs +++ b/src/value.rs @@ -34,6 +34,9 @@ pub enum Value<'lua> { Integer(Integer), /// A floating point number. Number(Number), + /// A Luau vector. + #[cfg(feature = "luau")] + Vector(f32, f32, f32), /// An interned string, managed by Lua. /// /// Unlike Rust strings, Lua strings may not be valid UTF-8. @@ -61,6 +64,8 @@ impl<'lua> Value<'lua> { Value::LightUserData(_) => "lightuserdata", Value::Integer(_) => "integer", Value::Number(_) => "number", + #[cfg(feature = "luau")] + Value::Vector(_, _, _) => "vector", Value::String(_) => "string", Value::Table(_) => "table", Value::Function(_) => "function", @@ -99,6 +104,8 @@ impl<'lua> PartialEq for Value<'lua> { (Value::Integer(a), Value::Number(b)) => *a as Number == *b, (Value::Number(a), Value::Integer(b)) => *a == *b as Number, (Value::Number(a), Value::Number(b)) => *a == *b, + #[cfg(feature = "luau")] + (Value::Vector(x1, y1, z1), Value::Vector(x2, y2, z2)) => (x1, y1, z1) == (x2, y2, z2), (Value::String(a), Value::String(b)) => a == b, (Value::Table(a), Value::Table(b)) => a == b, (Value::Function(a), Value::Function(b)) => a == b, @@ -130,6 +137,8 @@ impl<'lua> Serialize for Value<'lua> { .serialize_i64((*i).try_into().expect("cannot convert lua_Integer to i64")), #[allow(clippy::useless_conversion)] Value::Number(n) => serializer.serialize_f64(*n), + #[cfg(feature = "luau")] + Value::Vector(x, y, z) => (x, y, z).serialize(serializer), Value::String(s) => s.serialize(serializer), Value::Table(t) => t.serialize(serializer), Value::UserData(ud) => ud.serialize(serializer), diff --git a/tests/luau.rs b/tests/luau.rs index 2524efb..a32de00 100644 --- a/tests/luau.rs +++ b/tests/luau.rs @@ -3,7 +3,7 @@ use std::env; use std::fs; -use mlua::{Lua, Result}; +use mlua::{Lua, Result, Value}; #[test] fn test_require() -> Result<()> { @@ -29,3 +29,21 @@ fn test_require() -> Result<()> { ) .exec() } + +#[test] +fn test_vectors() -> Result<()> { + let lua = Lua::new(); + + let globals = lua.globals(); + globals.set( + "vector", + lua.create_function(|_, (x, y, z)| Ok(Value::Vector(x, y, z)))?, + )?; + + let v: [f32; 3] = lua + .load("return vector(1, 2, 3) + vector(3, 2, 1)") + .eval()?; + assert_eq!(v, [4.0, 4.0, 4.0]); + + Ok(()) +} diff --git a/tests/serde.rs b/tests/serde.rs index f5a8589..db0d673 100644 --- a/tests/serde.rs +++ b/tests/serde.rs @@ -144,6 +144,29 @@ fn test_serialize_failure() -> Result<(), Box> { Ok(()) } +#[cfg(feature = "luau")] +#[test] +fn test_serialize_vector() -> Result<(), Box> { + let lua = Lua::new(); + + let globals = lua.globals(); + globals.set( + "vector", + lua.create_function(|_, (x, y, z)| Ok(Value::Vector(x, y, z)))?, + )?; + + let val = lua.load("{_vector = vector(1, 2, 3)}").eval::()?; + let json = serde_json::json!({ + "_vector": [1.0, 2.0, 3.0], + }); + assert_eq!(serde_json::to_value(&val)?, json); + + let expected_json = lua.from_value::(val)?; + assert_eq!(expected_json, json); + + Ok(()) +} + #[test] fn test_to_value_struct() -> LuaResult<()> { let lua = Lua::new();