accessor_macro/
lib.rs

1use proc_macro::TokenStream;
2use quote::{format_ident, quote};
3use syn::{parse_macro_input, Data, DeriveInput, Fields};
4
5/// 为结构体派生 Getter 和 Setter 宏
6///
7/// 此宏会为带有 `accessor(get)` 属性的字段生成 getter 方法
8/// 为带有 `accessor(set)` 属性的字段生成 setter 方法。
9///
10/// 还支持通过 `accessor(range=[min, max])` 对字段设置值时进行范围检查。
11///
12/// # 示例
13///
14/// ```rust
15/// use accessor_macro::Accessor;
16///
17/// #[derive(Accessor, Debug)]
18/// struct Person {
19///    #[accessor(get, set)]
20///    name: String,
21///    #[accessor(get, set, range=[0, 200])]
22///    age: i32,
23/// }
24///
25/// let mut person = Person {
26///     name: "Alice".to_string(),
27///     age: 25,
28/// };
29///
30/// assert!(person.get_name().eq("Alice"));
31/// person.set_name("Bob".to_string());
32/// assert!(person.get_name().eq("Bob"));
33/// assert!(person.set_age(18));
34/// assert!(person.get_age() == &18);
35/// assert!(!person.set_age(-1));
36/// assert!(person.get_age() != &-1);
37/// assert!(!person.set_age(201));
38/// assert!(person.get_age() != &201);
39/// ```
40#[proc_macro_derive(Accessor, attributes(accessor))]
41pub fn accessor_derive(input: TokenStream) -> TokenStream {
42    let ast = parse_macro_input!(input as DeriveInput);
43    let name = &ast.ident;
44    let fields = if let Data::Struct(data_struct) = &ast.data {
45        if let Fields::Named(fields) = &data_struct.fields {
46            &fields.named
47        } else {
48            panic!("Only structs with named fields are supported");
49        }
50    } else {
51        panic!("Accessor can only be derived for structs");
52    };
53
54    let getters = fields.iter().filter_map(|field| {
55        let field_name = &field.ident;
56        let field_ty = &field.ty;
57        let attrs = &field.attrs;
58        let unaligned = attrs.iter().any(|attr| {
59            attr.path.is_ident("accessor") && attr.tokens.to_string().contains("unaligned")
60        });
61
62        if attrs
63            .iter()
64            .any(|attr| attr.path.is_ident("accessor") && attr.tokens.to_string().contains("get"))
65        {
66            let getter_name = format_ident!("get_{}", field_name.as_ref().unwrap());
67            if !unaligned {
68                Some(quote! {
69                    pub fn #getter_name(&self) -> &#field_ty {
70                        &self.#field_name
71                    }
72                })
73            } else {
74                Some(quote! {
75                    pub fn #getter_name(&self) -> #field_ty {
76                    let filed_ptr = std::ptr::addr_of!(self.#field_name);
77                        unsafe {
78                            filed_ptr.read_unaligned()
79                        }
80                    }
81                })
82            }
83        } else {
84            None
85        }
86    });
87
88    let setters = fields.iter().filter_map(|field| {
89        let field_name = &field.ident;
90        let field_ty = &field.ty;
91        let attrs = &field.attrs;
92
93        if attrs
94            .iter()
95            .any(|attr| attr.path.is_ident("accessor") && attr.tokens.to_string().contains("set"))
96        {
97            let setter_name = format_ident!("set_{}", field_name.as_ref().unwrap());
98            let range_check = attrs.iter().find_map(|attr| {
99                if attr.path.is_ident("accessor") {
100                    let tokens = attr.tokens.to_string();
101
102                    if let Some(range_start) = tokens.find("range=[") {
103                        let range_str = &tokens[range_start + 7..];
104                        let end_index = range_str.find(']').unwrap();
105                        let range_values = &range_str[0..end_index];
106                        let mut parts = range_values.split(',');
107                        let min = parts.next().unwrap().trim();
108                        let max = parts.next().unwrap().trim();
109
110                        let min_lit = syn::parse_str::<syn::Expr>(min).ok()?;
111                        let max_lit = syn::parse_str::<syn::Expr>(max).ok()?;
112
113                        #[cfg(all(debug_assertions, feature = "debug_panic"))]
114                        let out_of_range_handler = Some(quote!{
115                            panic!("field '{}' must be between {} and {}", stringify!(#field_name), #min_lit, #max_lit);
116                        });
117                        #[cfg(any(not(debug_assertions), not(feature = "debug_panic")))]
118                        let out_of_range_handler = Some(quote!{
119                            return false;
120                        });
121
122                        Some(quote! {
123                            if value < #min_lit || value > #max_lit {
124                                #out_of_range_handler
125                            }
126                        })
127                    } else {
128                        None
129                    }
130                } else {
131                    None
132                }
133            });
134
135            Some(quote! {
136                pub fn #setter_name(&mut self, value: #field_ty) -> bool {
137                    #range_check
138                    self.#field_name = value;
139                    true
140                }
141            })
142        } else {
143            None
144        }
145    });
146
147    let expanded = quote! {
148        impl #name {
149            #(#getters)*
150            #(#setters)*
151        }
152    };
153
154    TokenStream::from(expanded)
155}