1use proc_macro::TokenStream;
2use quote::{format_ident, quote};
3use syn::{parse_macro_input, Data, DeriveInput, Fields};
4
5#[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 = ¶m_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}