Skip to main content

get_size_derive2/
lib.rs

1//! Derives the `GetSize` trait of [`get-size2`](https://docs.rs/get-size2) for structs and enums.
2//!
3//! This crate is re-exported by `get-size2` and is meant to be used through its `derive` feature.
4//! See [`GetSize`](macro@GetSize) for the attribute reference.
5
6use 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    // Skip all attributes which do not belong to us.
37    if !attr.meta.path().is_ident("get_size") {
38        return collection;
39    }
40
41    // Make sure it is a list: #[get_size(...)]
42    let Ok(list) = attr.meta.require_list() else {
43        return collection;
44    };
45
46    // Parse the nested meta: #[get_size(ignore(...))] or #[get_size(ignore)]
47    let _ = list.parse_nested_meta(|meta| {
48        // Only handle `ignore`
49        if !meta.path.is_ident("ignore") {
50            return Ok(()); // Skip unrelated
51        }
52
53        // Handle the flag case: #[get_size(ignore)]
54        if meta.input.is_empty() {
55            // Do nothing – valid empty ignore
56            return Ok(());
57        }
58
59        // Handle the list case: #[get_size(ignore(A, B))]
60        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
96// Add a bound `T: GetSize` to every type parameter T, unless we ignore it.
97fn 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    // Construct a representation of Rust code as a syntax tree that we can manipulate.
132    let ast: syn::DeriveInput = syn::parse(input)?;
133
134    // The name of the struct.
135    let name = &ast.ident;
136
137    // Extract all generics we shall ignore.
138    // let ignored = extract_ignored_generics_list(&ast.attrs);
139    let ignored = collect_all_ignored_generics(&ast);
140
141    // Add a bound `T: GetSize` to every type parameter T.
142    let generics = add_trait_bounds(ast.generics, &ignored);
143
144    // Extract the generics of the struct/enum.
145    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
146
147    // Traverse the parsed data to generate the individual parts of the function.
148    match ast.data {
149        syn::Data::Enum(data_enum) => {
150            if data_enum.variants.is_empty() {
151                // Empty enums are easy to implement.
152                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                            // Parse all relevant attributes.
172                            let attr = StructFieldAttribute::from_attributes(&field.attrs)
173                                .map_err(|err| syn::Error::new_spanned(field, err.to_string()))?;
174
175                            // Fields handled by `size` or `ignore` are never read, so they are
176                            // bound to a wildcard to avoid unused variable warnings.
177                            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                            // Fields handled by `size` or `ignore` are never read, so they stay
230                            // out of the pattern and are covered by its `..` rest pattern.
231                            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            // Build the trait implementation
282            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                // Empty structs are easy to implement.
311                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 newtypes
320
321            for field in &data_struct.fields {
322                // Parse all relevant attributes.
323                let attr = StructFieldAttribute::from_attributes(&field.attrs)
324                    .map_err(|err| syn::Error::new_spanned(field, err.to_string()))?;
325
326                // How this field is accessed: by name, or by position for tuple structs. The
327                // positional counter has to advance even when the field is handled by one of the
328                // attributes below, otherwise all following fields of a tuple struct would be
329                // read at the wrong position.
330                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            // Build the trait implementation
362            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}