diff --git a/mlua_derive/src/from_lua.rs b/mlua_derive/src/from_lua.rs new file mode 100644 index 0000000..9e311e4 --- /dev/null +++ b/mlua_derive/src/from_lua.rs @@ -0,0 +1,32 @@ +use proc_macro::TokenStream; +use quote::quote; +use syn::{parse_macro_input, DeriveInput}; + +pub fn from_lua(input: TokenStream) -> TokenStream { + let DeriveInput { + ident, generics, .. + } = parse_macro_input!(input as DeriveInput); + + let where_clause = match &generics.where_clause { + Some(where_clause) => quote! { #where_clause, Self: 'static + Clone }, + None => quote! { where Self: 'static + Clone }, + }; + let ident_str = ident.to_string(); + + quote! { + impl #generics ::mlua::FromLua<'_> for #ident #generics #where_clause { + #[inline] + fn from_lua(value: ::mlua::Value<'_>, lua: &'_ ::mlua::Lua) -> ::mlua::Result { + match value { + ::mlua::Value::UserData(ud) => Ok(ud.borrow::()?.clone()), + _ => Err(::mlua::Error::FromLuaConversionError { + from: value.type_name(), + to: #ident_str, + message: None, + }), + } + } + } + } + .into() +} diff --git a/mlua_derive/src/lib.rs b/mlua_derive/src/lib.rs index 59f77a9..78937a8 100644 --- a/mlua_derive/src/lib.rs +++ b/mlua_derive/src/lib.rs @@ -148,7 +148,15 @@ pub fn chunk(input: TokenStream) -> TokenStream { wrapped_code.into() } +#[cfg(feature = "macros")] +#[proc_macro_derive(FromLua)] +pub fn from_lua(input: TokenStream) -> TokenStream { + from_lua::from_lua(input) +} + #[cfg(feature = "macros")] mod chunk; #[cfg(feature = "macros")] +mod from_lua; +#[cfg(feature = "macros")] mod token; diff --git a/src/lib.rs b/src/lib.rs index 26d5fef..4f6ea09 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -216,6 +216,14 @@ pub use crate::{ #[cfg_attr(docsrs, doc(cfg(feature = "macros")))] pub use mlua_derive::chunk; +/// Derive [`FromLua`] for a Rust type. +/// +/// Current implementation generate code that takes [`UserData`] value, borrow it (of the Rust type) +/// and clone. +#[cfg(feature = "macros")] +#[cfg_attr(docsrs, doc(cfg(feature = "macros")))] +pub use mlua_derive::FromLua; + /// Registers Lua module entrypoint. /// /// You can register multiple entrypoints as required. diff --git a/tests/userdata.rs b/tests/userdata.rs index 1be4b22..8bdd46b 100644 --- a/tests/userdata.rs +++ b/tests/userdata.rs @@ -908,3 +908,39 @@ fn test_owned_userdata() -> Result<()> { Ok(()) } + +#[cfg(feature = "macros")] +#[test] +fn test_userdata_derive() -> Result<()> { + let lua = Lua::new(); + + // Simple struct + + #[derive(Clone, Copy, mlua::FromLua)] + struct MyUserData(i32); + + lua.register_userdata_type::(|reg| { + reg.add_function("val", |_, this: MyUserData| Ok(this.0)); + })?; + + lua.globals() + .set("ud", AnyUserData::wrap(MyUserData(123)))?; + lua.load("assert(ud:val() == 123)").exec()?; + + // More complex struct where generics and where clause + + #[derive(Clone, Copy, mlua::FromLua)] + struct MyUserData2<'a, T>(&'a T) + where + T: ?Sized; + + lua.register_userdata_type::>(|reg| { + reg.add_function("val", |_, this: MyUserData2<'static, i32>| Ok(*this.0)); + })?; + + lua.globals() + .set("ud", AnyUserData::wrap(MyUserData2(&321)))?; + lua.load("assert(ud:val() == 321)").exec()?; + + Ok(()) +}