Skip to main content

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