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        let mut no_ref = false;
63
64        if attrs.iter().any(|attr| {
65            let param_string = attr.tokens.to_string();
66            if let Some(param_start) = param_string.find("get(") {
67                if let Some(param_end) = param_string[param_start..].find(')') {
68                    let params = &param_string[param_start + 4..=param_end];
69                    if params.contains("no_ref") {
70                        no_ref = true;
71                    }
72                }
73            }
74            attr.path.is_ident("accessor") && attr.tokens.to_string().contains("get")
75        }) {
76            let getter_name = format_ident!("get_{}", field_name.as_ref().unwrap());
77            if unaligned {
78                Some(quote! {
79                    pub fn #getter_name(&self) -> #field_ty {
80                    let filed_ptr = std::ptr::addr_of!(self.#field_name);
81                        unsafe {
82                            filed_ptr.read_unaligned()
83                        }
84                    }
85                })
86            } else if no_ref {
87                Some(quote! {
88                    pub fn #getter_name(&self) -> #field_ty {
89                        self.#field_name
90                    }
91                })
92            } else {
93                Some(quote! {
94                    pub fn #getter_name(&self) -> &#field_ty {
95                        &self.#field_name
96                    }
97                })
98            }
99        } else {
100            None
101        }
102    });
103
104    let setters = fields.iter().filter_map(|field| {
105        let field_name = &field.ident;
106        let field_ty = &field.ty;
107        let attrs = &field.attrs;
108
109        if attrs
110            .iter()
111            .any(|attr| attr.path.is_ident("accessor") && attr.tokens.to_string().contains("set"))
112        {
113            let setter_name = format_ident!("set_{}", field_name.as_ref().unwrap());
114            let range_check = attrs.iter().find_map(|attr| {
115                if attr.path.is_ident("accessor") {
116                    let tokens = attr.tokens.to_string();
117
118                    if let Some(range_start) = tokens.find("range=[") {
119                        let range_str = &tokens[range_start + 7..];
120                        let end_index = range_str.find(']').unwrap();
121                        let range_values = &range_str[0..end_index];
122                        let mut parts = range_values.split(',');
123                        let min = parts.next().unwrap().trim();
124                        let max = parts.next().unwrap().trim();
125
126                        let min_lit = syn::parse_str::<syn::Expr>(min).ok()?;
127                        let max_lit = syn::parse_str::<syn::Expr>(max).ok()?;
128
129                        #[cfg(all(debug_assertions, feature = "debug_panic"))]
130                        let out_of_range_handler = Some(quote!{
131                            panic!("field '{}' must be between {} and {}", stringify!(#field_name), #min_lit, #max_lit);
132                        });
133                        #[cfg(any(not(debug_assertions), not(feature = "debug_panic")))]
134                        let out_of_range_handler = Some(quote!{
135                            return false;
136                        });
137
138                        Some(quote! {
139                            if value < #min_lit || value > #max_lit {
140                                #out_of_range_handler
141                            }
142                        })
143                    } else {
144                        None
145                    }
146                } else {
147                    None
148                }
149            });
150
151            Some(quote! {
152                pub fn #setter_name(&mut self, value: #field_ty) -> bool {
153                    #range_check
154                    self.#field_name = value;
155                    true
156                }
157            })
158        } else {
159            None
160        }
161    });
162
163    let expanded = quote! {
164        impl #name {
165            #(#getters)*
166            #(#setters)*
167        }
168    };
169
170    TokenStream::from(expanded)
171}