1use attribute_derive::{Attribute, FromAttr};
7use proc_macro::TokenStream;
8use quote::{format_ident, quote};
9
10#[derive(FromAttr, Default, Debug)]
11#[attribute(ident = get_size)]
12struct StructFieldAttribute {
13 #[attribute(conflicts = [size_fn, ignore])]
14 size: Option<usize>,
15 #[attribute(conflicts = [size, ignore])]
16 size_fn: Option<syn::Ident>,
17 #[attribute(conflicts = [size, size_fn])]
18 ignore: bool,
19}
20
21fn extract_ignored_generics_list(list: &Vec<syn::Attribute>) -> Vec<syn::PathSegment> {
22 let mut collection = Vec::new();
23
24 for attr in list {
25 let mut list = extract_ignored_generics(attr);
26
27 collection.append(&mut list);
28 }
29
30 collection
31}
32
33fn extract_ignored_generics(attr: &syn::Attribute) -> Vec<syn::PathSegment> {
34 let mut collection = Vec::new();
35
36 if !attr.meta.path().is_ident("get_size") {
38 return collection;
39 }
40
41 let Ok(list) = attr.meta.require_list() else {
43 return collection;
44 };
45
46 let _ = list.parse_nested_meta(|meta| {
48 if !meta.path.is_ident("ignore") {
50 return Ok(()); }
52
53 if meta.input.is_empty() {
55 return Ok(());
57 }
58
59 meta.parse_nested_meta(|meta| {
61 for segment in meta.path.segments {
62 collection.push(segment);
63 }
64 Ok(())
65 })?;
66
67 Ok(())
68 });
69
70 collection
71}
72
73fn collect_all_ignored_generics(ast: &syn::DeriveInput) -> Vec<syn::PathSegment> {
74 let mut ignored = extract_ignored_generics_list(&ast.attrs);
75
76 match &ast.data {
77 syn::Data::Struct(data_struct) => {
78 for field in &data_struct.fields {
79 ignored.extend(extract_ignored_generics_list(&field.attrs));
80 }
81 }
82 syn::Data::Enum(data_enum) => {
83 for variant in &data_enum.variants {
84 ignored.extend(extract_ignored_generics_list(&variant.attrs));
85 for field in &variant.fields {
86 ignored.extend(extract_ignored_generics_list(&field.attrs));
87 }
88 }
89 }
90 syn::Data::Union(_) => {}
91 }
92
93 ignored
94}
95
96fn add_trait_bounds(mut generics: syn::Generics, ignored: &Vec<syn::PathSegment>) -> syn::Generics {
98 for param in &mut generics.params {
99 if let syn::GenericParam::Type(type_param) = param {
100 let mut found = false;
101 for ignored in ignored {
102 if ignored.ident == type_param.ident {
103 found = true;
104 break;
105 }
106 }
107
108 if found {
109 continue;
110 }
111
112 type_param
113 .bounds
114 .push(syn::parse_quote!(::get_size2::GetSize));
115 }
116 }
117 generics
118}
119
120#[doc = include_str!("./derive.md")]
121#[proc_macro_derive(GetSize, attributes(get_size))]
122pub fn derive_get_size(input: TokenStream) -> TokenStream {
123 match derive_get_size_impl(input) {
124 Ok(tokens) => tokens,
125 Err(err) => err.to_compile_error().into(),
126 }
127}
128
129#[expect(clippy::too_many_lines, reason = "Needs refactoring")]
130fn derive_get_size_impl(input: TokenStream) -> syn::Result<TokenStream> {
131 let ast: syn::DeriveInput = syn::parse(input)?;
133
134 let name = &ast.ident;
136
137 let ignored = collect_all_ignored_generics(&ast);
140
141 let generics = add_trait_bounds(ast.generics, &ignored);
143
144 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
146
147 match ast.data {
149 syn::Data::Enum(data_enum) => {
150 if data_enum.variants.is_empty() {
151 let generated = quote! {
153 impl ::get_size2::GetSize for #name {}
154 };
155 return Ok(generated.into());
156 }
157
158 let mut cmds = Vec::with_capacity(data_enum.variants.len());
159
160 for variant in data_enum.variants {
161 let ident = &variant.ident;
162
163 match &variant.fields {
164 syn::Fields::Unnamed(unnamed_fields) => {
165 let num_fields = unnamed_fields.unnamed.len();
166
167 let mut field_idents = Vec::with_capacity(num_fields);
168 let mut field_cmds = Vec::with_capacity(num_fields);
169
170 for (i, field) in unnamed_fields.unnamed.iter().enumerate() {
171 let attr = StructFieldAttribute::from_attributes(&field.attrs)
173 .map_err(|err| syn::Error::new_spanned(field, err.to_string()))?;
174
175 if let Some(size) = attr.size {
178 field_idents.push(quote! { _ });
179 field_cmds.push(quote! {
180 total += #size;
181 });
182
183 continue;
184 } else if attr.ignore {
185 field_idents.push(quote! { _ });
186
187 continue;
188 }
189
190 let field_ident = format_ident!("v{i}");
191
192 if let Some(size_fn) = attr.size_fn {
193 field_cmds.push(quote! {
194 total += #size_fn(#field_ident);
195 });
196 } else {
197 field_cmds.push(quote! {
198 let (total_add, tracker) = ::get_size2::GetSize::get_heap_size_with_tracker(#field_ident, tracker);
199 total += total_add;
200 });
201 }
202
203 field_idents.push(quote! { #field_ident });
204 }
205
206 cmds.push(quote! {
207 Self::#ident(#(#field_idents,)*) => {
208 let mut total = 0;
209
210 #(#field_cmds)*;
211
212 (total, tracker)
213 }
214 });
215 }
216 syn::Fields::Named(named_fields) => {
217 let mut field_idents = Vec::new();
218 let mut field_cmds = Vec::new();
219 let mut skipped_field = false;
220
221 for field in &named_fields.named {
222 let field_ident = field.ident.as_ref().ok_or_else(|| {
223 syn::Error::new_spanned(field, "Expected named field")
224 })?;
225
226 let attr = StructFieldAttribute::from_attributes(&field.attrs)
227 .map_err(|err| syn::Error::new_spanned(field, err.to_string()))?;
228
229 if let Some(size) = attr.size {
232 skipped_field = true;
233 field_cmds.push(quote! {
234 total += #size;
235 });
236
237 continue;
238 } else if attr.ignore {
239 skipped_field = true;
240
241 continue;
242 }
243
244 field_idents.push(field_ident);
245
246 if let Some(size_fn) = attr.size_fn {
247 field_cmds.push(quote! {
248 total += #size_fn(#field_ident);
249 });
250 } else {
251 field_cmds.push(quote! {
252 let (total_add, tracker) = ::get_size2::GetSize::get_heap_size_with_tracker(#field_ident, tracker);
253 total += total_add;
254 });
255 }
256 }
257
258 let pattern = if skipped_field {
259 quote! { Self::#ident { #(#field_idents,)* .. } }
260 } else {
261 quote! { Self::#ident { #(#field_idents,)* } }
262 };
263
264 cmds.push(quote! {
265 #pattern => {
266 let mut total = 0;
267 #(#field_cmds)*
268 (total, tracker)
269 }
270 });
271 }
272
273 syn::Fields::Unit => {
274 cmds.push(quote! {
275 Self::#ident => (0, tracker),
276 });
277 }
278 }
279 }
280
281 let generated = quote! {
283 impl #impl_generics ::get_size2::GetSize for #name #ty_generics #where_clause {
284 fn get_heap_size(&self) -> usize {
285 let tracker = ::get_size2::default_tracker();
286
287 let (total, _) = ::get_size2::GetSize::get_heap_size_with_tracker(self, tracker);
288
289 total
290 }
291
292 fn get_heap_size_with_tracker<TRACKER: ::get_size2::GetSizeTracker>(
293 &self,
294 tracker: TRACKER,
295 ) -> (usize, TRACKER) {
296 match self {
297 #(#cmds)*
298 }
299 }
300 }
301 };
302 Ok(generated.into())
303 }
304 syn::Data::Union(_data_union) => Err(syn::Error::new_spanned(
305 name,
306 "Deriving GetSize for unions is currently not supported.",
307 )),
308 syn::Data::Struct(data_struct) => {
309 if data_struct.fields.is_empty() {
310 let generated = quote! {
312 impl ::get_size2::GetSize for #name {}
313 };
314 return Ok(generated.into());
315 }
316
317 let mut cmds = Vec::with_capacity(data_struct.fields.len());
318
319 let mut unidentified_fields_count = 0; for field in &data_struct.fields {
322 let attr = StructFieldAttribute::from_attributes(&field.attrs)
324 .map_err(|err| syn::Error::new_spanned(field, err.to_string()))?;
325
326 let accessor = field.ident.as_ref().map_or_else(
331 || {
332 let index = syn::Index::from(unidentified_fields_count);
333 unidentified_fields_count += 1;
334 quote! { #index }
335 },
336 |ident| quote! { #ident },
337 );
338
339 if let Some(size) = attr.size {
340 cmds.push(quote! {
341 total += #size;
342 });
343
344 continue;
345 } else if let Some(size_fn) = attr.size_fn {
346 cmds.push(quote! {
347 total += #size_fn(&self.#accessor);
348 });
349
350 continue;
351 } else if attr.ignore {
352 continue;
353 }
354
355 cmds.push(quote! {
356 let (total_add, tracker) = ::get_size2::GetSize::get_heap_size_with_tracker(&self.#accessor, tracker);
357 total += total_add;
358 });
359 }
360
361 let generated = quote! {
363 impl #impl_generics ::get_size2::GetSize for #name #ty_generics #where_clause {
364 fn get_heap_size(&self) -> usize {
365 let tracker = ::get_size2::default_tracker();
366
367 let (total, _) = ::get_size2::GetSize::get_heap_size_with_tracker(self, tracker);
368
369 total
370 }
371
372 fn get_heap_size_with_tracker<TRACKER: ::get_size2::GetSizeTracker>(
373 &self,
374 tracker: TRACKER,
375 ) -> (usize, TRACKER) {
376 let mut total = 0;
377
378 #(#cmds)*;
379
380 (total, tracker)
381 }
382 }
383 };
384 Ok(generated.into())
385 }
386 }
387}