diff --git a/Cargo.toml b/Cargo.toml index be3b04b..751f0af 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "osui-element" -version = "0.1.9" +version = "0.1.8" edition = "2021" description = "The element attribute for defining elements in OSUI" license = "Apache-2.0" @@ -9,6 +9,10 @@ license = "Apache-2.0" proc-macro = true path = "./lib.rs" +[[bin]] +name = "osui-test" +path = "./main.rs" + [dependencies] proc-macro2 = "1.0.89" quote = "1.0" diff --git a/lib.rs b/lib.rs index fc6ce96..96b9059 100644 --- a/lib.rs +++ b/lib.rs @@ -7,9 +7,13 @@ use syn::{ #[proc_macro_attribute] pub fn element(_args: TokenStream, input: TokenStream) -> TokenStream { let mut ast = parse_macro_input!(input as DeriveInput); + let struct_name = &ast.ident; - // Add a lifetime if none exists - if ast.generics.lifetimes().count() == 0 { + // Check if the struct has generics + let has_generics = !ast.generics.params.is_empty(); + + // If there are no generics, add a lifetime (if necessary) + if ast.generics.lifetimes().count() == 0 && !has_generics { ast.generics .params .push(GenericParam::Lifetime(LifetimeParam::new(Lifetime::new( @@ -47,8 +51,90 @@ pub fn element(_args: TokenStream, input: TokenStream) -> TokenStream { _ => panic!("`element` can only be used with structs"), } + // Generate the impl block with generics if needed + let impl_block = if has_generics { + // Keep the struct's generics (if any) when implementing the trait + quote! { + impl<'a, T> ElementCore for #struct_name<'a, T> { + fn get_element_by_id(&mut self, id: &str) -> Option<&mut Element> { + if let Children::Children(children, _) = &mut self.children { + for elem in children { + if elem.get_id() == id { + // return Some(elem); + } else if let Some(e) = elem.get_element_by_id(id) { + return Some(e); + } + } + } + None + } + + fn set_styling(&mut self, styling: &std::collections::HashMap) { + if let Some(style) = styling.get(&StyleName::Class(self.class.to_string())) { + self.style = style.clone(); + } else if let Some(style) = styling.get(&StyleName::Id(self.id.to_string())) { + self.style = style.clone(); + } else if let Some(style) = + styling.get(&StyleName::Component(stringify!(#struct_name).to_string())) + { + self.style = style.clone(); + } + if let Children::Children(children, _) = &mut self.children { + for child in children { + child.set_styling(styling); + } + } + } + + fn get_id(&self) -> String { + self.id.to_string() + } + } + } + } else { + // If no generics, don't add a generic `T` to the trait implementation + quote! { + impl<'a> ElementCore for #struct_name<'a> { + fn get_element_by_id(&mut self, id: &str) -> Option<&mut Element> { + if let Children::Children(children, _) = &mut self.children { + for elem in children { + if elem.get_id() == id { + // return Some(elem); + } else if let Some(e) = elem.get_element_by_id(id) { + return Some(e); + } + } + } + None + } + + fn set_styling(&mut self, styling: &std::collections::HashMap) { + if let Some(style) = styling.get(&StyleName::Class(self.class.to_string())) { + self.style = style.clone(); + } else if let Some(style) = styling.get(&StyleName::Id(self.id.to_string())) { + self.style = style.clone(); + } else if let Some(style) = + styling.get(&StyleName::Component(stringify!(#struct_name).to_string())) + { + self.style = style.clone(); + } + if let Children::Children(children, _) = &mut self.children { + for child in children { + child.set_styling(styling); + } + } + } + + fn get_id(&self) -> String { + self.id.to_string() + } + } + } + }; + let expanded = quote! { #ast + #impl_block }; expanded.into() @@ -58,17 +144,27 @@ pub fn element(_args: TokenStream, input: TokenStream) -> TokenStream { pub fn elem_fn(_args: TokenStream, input: TokenStream) -> TokenStream { let ast = parse_macro_input!(input as DeriveInput); let struct_name = &ast.ident; + let func_name = syn::Ident::new( &struct_name.to_string().to_lowercase(), proc_macro2::Span::call_site(), ); - let elem_fn = quote! { - pub fn #func_name<'a, T>() -> Box<#struct_name<'a, T>> { - Box::new(#struct_name::default()) + let elem_fn = if ast.generics.params.len() == 1 { + quote! { + pub fn #func_name<'a>() -> Box<#struct_name<'a>> { + Box::new(#struct_name::default()) + } + } + } else { + quote! { + pub fn #func_name<'a, T>() -> Box<#struct_name<'a, T>> { + Box::new(#struct_name::default()) + } } }; + // Combine the struct and the generated function let expanded = quote! { #ast #elem_fn