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 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}