Skip to main content

serde_derive/
bound.rs

1use crate::internals::ast::{Container, Data};
2use crate::internals::{attr, ungroup};
3use proc_macro2::Span;
4use std::collections::HashSet;
5use syn::punctuated::{Pair, Punctuated};
6use syn::Token;
7
8// Remove the default from every type parameter because in the generated impls
9// they look like associated types: "error: associated type bindings are not
10// allowed here".
11pub fn without_defaults(generics: &syn::Generics) -> syn::Generics {
12    syn::Generics {
13        params: generics
14            .params
15            .iter()
16            .map(|param| match param {
17                syn::GenericParam::Type(param) => syn::GenericParam::Type(syn::TypeParam {
18                    default: None,
19                    ..param.clone()
20                }),
21                _ => param.clone(),
22            })
23            .collect(),
24        ..generics.clone()
25    }
26}
27
28pub fn with_where_predicates(
29    generics: &syn::Generics,
30    predicates: &[syn::WherePredicate],
31) -> syn::Generics {
32    let mut generics = generics.clone();
33    let dst_predicates = &mut generics.make_where_clause().predicates;
34
35    for predicate in predicates {
36        dst_predicates.push(predicate.clone());
37    }
38    generics
39}
40
41pub fn with_where_predicates_from_fields(
42    cont: &Container,
43    generics: &syn::Generics,
44    from_field: fn(&attr::Field) -> Option<&[syn::WherePredicate]>,
45) -> syn::Generics {
46    let mut generics = generics.clone();
47    let dst_predicates = &mut generics.make_where_clause().predicates;
48
49    for field in cont.data.all_fields() {
50        let Some(predicate_slice) = from_field(&field.attrs) else {
51            continue;
52        };
53        for inner_predicate in predicate_slice {
54            dst_predicates.push(inner_predicate.clone());
55        }
56    }
57    generics
58}
59
60pub fn with_where_predicates_from_variants(
61    cont: &Container,
62    generics: &syn::Generics,
63    from_variant: fn(&attr::Variant) -> Option<&[syn::WherePredicate]>,
64) -> syn::Generics {
65    let variants = match &cont.data {
66        Data::Enum(variants) => variants,
67        Data::Struct(_, _) => {
68            return generics.clone();
69        }
70    };
71    let mut generics = generics.clone();
72    let dst_predicates = &mut generics.make_where_clause().predicates;
73
74    for variant in variants {
75        let Some(predicate_slice) = from_variant(&variant.attrs) else {
76            continue;
77        };
78        for inner_predicate in predicate_slice {
79            dst_predicates.push(inner_predicate.clone());
80        }
81    }
82    generics
83}
84
85// Puts the given bound on any generic type parameters that are used in fields
86// for which filter returns true.
87//
88// For example, the following struct needs the bound `A: Serialize, B:
89// Serialize`.
90//
91//     struct S<'b, A, B: 'b, C> {
92//         a: A,
93//         b: Option<&'b B>
94//         #[serde(skip_serializing)]
95//         c: C,
96//     }
97pub fn with_bound(
98    cont: &Container,
99    generics: &syn::Generics,
100    filter: fn(&attr::Field, Option<&attr::Variant>) -> bool,
101    bound: &syn::Path,
102) -> syn::Generics {
103    struct FindTyParams<'ast> {
104        // Set of all generic type parameters on the current struct (A, B, C in
105        // the example). Initialized up front.
106        all_type_params: HashSet<syn::Ident>,
107
108        // Set of generic type parameters used in fields for which filter
109        // returns true (A and B in the example). Filled in as the visitor sees
110        // them.
111        relevant_type_params: HashSet<syn::Ident>,
112
113        // Fields whose type is an associated type of one of the generic type
114        // parameters.
115        associated_type_usage: Vec<&'ast syn::TypePath>,
116    }
117
118    impl<'ast> FindTyParams<'ast> {
119        fn visit_field(&mut self, field: &'ast syn::Field) {
120            if let syn::Type::Path(ty) = ungroup(&field.ty) {
121                if let Some(Pair::Punctuated(t, _)) = ty.path.segments.pairs().next() {
122                    if self.all_type_params.contains(&t.ident) {
123                        self.associated_type_usage.push(ty);
124                    }
125                }
126            }
127            self.visit_type(&field.ty);
128        }
129
130        fn visit_path(&mut self, path: &'ast syn::Path) {
131            if let Some(seg) = path.segments.last() {
132                if seg.ident == "PhantomData" {
133                    // Hardcoded exception, because PhantomData<T> implements
134                    // Serialize and Deserialize whether or not T implements it.
135                    return;
136                }
137            }
138            if path.leading_colon.is_none() && path.segments.len() == 1 {
139                let id = &path.segments[0].ident;
140                if self.all_type_params.contains(id) {
141                    self.relevant_type_params.insert(id.clone());
142                }
143            }
144            for segment in &path.segments {
145                self.visit_path_segment(segment);
146            }
147        }
148
149        // Everything below is simply traversing the syntax tree.
150
151        fn visit_type(&mut self, ty: &'ast syn::Type) {
152            match ty {
153                #![cfg_attr(all(test, exhaustive), deny(non_exhaustive_omitted_patterns))]
154                syn::Type::Array(ty) => self.visit_type(&ty.elem),
155                syn::Type::FnPtr(ty) => {
156                    for arg in &ty.inputs {
157                        self.visit_type(&arg.ty);
158                    }
159                    self.visit_return_type(&ty.output);
160                }
161                syn::Type::Group(ty) => self.visit_type(&ty.elem),
162                syn::Type::ImplTrait(ty) => {
163                    for bound in &ty.bounds {
164                        self.visit_type_param_bound(bound);
165                    }
166                }
167                syn::Type::Macro(ty) => self.visit_macro(&ty.mac),
168                syn::Type::Paren(ty) => self.visit_type(&ty.elem),
169                syn::Type::Path(ty) => {
170                    if let Some(qself) = &ty.qself {
171                        self.visit_type(&qself.ty);
172                    }
173                    self.visit_path(&ty.path);
174                }
175                syn::Type::Ptr(ty) => self.visit_type(&ty.elem),
176                syn::Type::Reference(ty) => self.visit_type(&ty.elem),
177                syn::Type::Slice(ty) => self.visit_type(&ty.elem),
178                syn::Type::TraitObject(ty) => {
179                    for bound in &ty.bounds {
180                        self.visit_type_param_bound(bound);
181                    }
182                }
183                syn::Type::Tuple(ty) => {
184                    for elem in &ty.elems {
185                        self.visit_type(elem);
186                    }
187                }
188
189                syn::Type::Infer(_) | syn::Type::Never(_) | syn::Type::Verbatim(_) => {}
190
191                _ => {}
192            }
193        }
194
195        fn visit_path_segment(&mut self, segment: &'ast syn::PathSegment) {
196            self.visit_path_arguments(&segment.arguments);
197        }
198
199        fn visit_path_arguments(&mut self, arguments: &'ast syn::PathArguments) {
200            match arguments {
201                syn::PathArguments::None => {}
202                syn::PathArguments::AngleBracketed(arguments) => {
203                    for arg in &arguments.args {
204                        match arg {
205                            #![cfg_attr(all(test, exhaustive), deny(non_exhaustive_omitted_patterns))]
206                            syn::GenericArgument::Type(arg) => self.visit_type(arg),
207                            syn::GenericArgument::AssocType(arg) => self.visit_type(&arg.ty),
208                            syn::GenericArgument::Lifetime(_)
209                            | syn::GenericArgument::Const(_)
210                            | syn::GenericArgument::AssocConst(_)
211                            | syn::GenericArgument::Constraint(_) => {}
212                            _ => {}
213                        }
214                    }
215                }
216                syn::PathArguments::Parenthesized(arguments) => {
217                    for argument in &arguments.inputs {
218                        self.visit_type(&argument.ty);
219                    }
220                    self.visit_return_type(&arguments.output);
221                }
222            }
223        }
224
225        fn visit_return_type(&mut self, return_type: &'ast syn::ReturnType) {
226            match return_type {
227                syn::ReturnType::Default => {}
228                syn::ReturnType::Type(_, output) => self.visit_type(output),
229            }
230        }
231
232        fn visit_type_param_bound(&mut self, bound: &'ast syn::TypeParamBound) {
233            match bound {
234                #![cfg_attr(all(test, exhaustive), deny(non_exhaustive_omitted_patterns))]
235                syn::TypeParamBound::Trait(bound) => self.visit_path(&bound.path),
236                syn::TypeParamBound::Lifetime(_)
237                | syn::TypeParamBound::PreciseCapture(_)
238                | syn::TypeParamBound::Verbatim(_) => {}
239                _ => {}
240            }
241        }
242
243        // Type parameter should not be considered used by a macro path.
244        //
245        //     struct TypeMacro<T> {
246        //         mac: T!(),
247        //         marker: PhantomData<T>,
248        //     }
249        fn visit_macro(&mut self, _mac: &'ast syn::Macro) {}
250    }
251
252    let all_type_params = generics
253        .type_params()
254        .map(|param| param.ident.clone())
255        .collect();
256
257    let mut visitor = FindTyParams {
258        all_type_params,
259        relevant_type_params: HashSet::new(),
260        associated_type_usage: Vec::new(),
261    };
262    match &cont.data {
263        Data::Enum(variants) => {
264            for variant in variants {
265                for field in &variant.fields {
266                    if filter(&field.attrs, Some(&variant.attrs)) {
267                        visitor.visit_field(field.original);
268                    }
269                }
270            }
271        }
272        Data::Struct(_, fields) => {
273            for field in fields {
274                if filter(&field.attrs, None) {
275                    visitor.visit_field(field.original);
276                }
277            }
278        }
279    }
280
281    let relevant_type_params = visitor.relevant_type_params;
282    let associated_type_usage = visitor.associated_type_usage;
283
284    fn make_where_bounded_type(
285        bounded_ty: syn::TypePath,
286        bound: &syn::Path,
287    ) -> syn::WherePredicate {
288        syn::WherePredicate::Type(syn::PredicateType {
289            attrs: Vec::new(),
290            lifetimes: None,
291            // the type parameter that is being bounded e.g. T
292            bounded_ty: syn::Type::Path(bounded_ty),
293            colon_token: <::syn::token::ColonToken![:]>::default(),
294            // the bound e.g. Serialize
295            bounds: {
296                let mut punct = Punctuated::new();
297                punct.push(syn::TypeParamBound::Trait(syn::TraitBound {
298                    paren_token: None,
299                    lifetimes: None,
300                    modifiers: syn::TraitBoundModifiers::default(),
301                    maybe: None,
302                    path: bound.clone(),
303                }));
304                punct
305            },
306        })
307    }
308
309    let mut dst_generics = generics.clone();
310    let dst_predicates = &mut dst_generics.make_where_clause().predicates;
311    for param in generics.type_params() {
312        let id = &param.ident;
313        if !relevant_type_params.contains(id) {
314            continue;
315        }
316        let bounded_ty = syn::TypePath {
317            attrs: Vec::new(),
318            qself: None,
319            path: id.clone().into(),
320        };
321        dst_predicates.push(make_where_bounded_type(bounded_ty, bound));
322    }
323    for bounded_ty in associated_type_usage {
324        dst_predicates.push(make_where_bounded_type(bounded_ty.clone(), bound));
325    }
326    dst_generics
327}
328
329pub fn with_self_bound(
330    cont: &Container,
331    generics: &syn::Generics,
332    bound: &syn::Path,
333) -> syn::Generics {
334    let mut generics = generics.clone();
335    generics
336        .make_where_clause()
337        .predicates
338        .push(syn::WherePredicate::Type(syn::PredicateType {
339            attrs: Vec::new(),
340            lifetimes: None,
341            // the type that is being bounded e.g. MyStruct<'a, T>
342            bounded_ty: type_of_item(cont),
343            colon_token: <::syn::token::ColonToken![:]>::default(),
344            // the bound e.g. Default
345            bounds: ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [syn::TypeParamBound::Trait(syn::TraitBound {
                        paren_token: None,
                        lifetimes: None,
                        modifiers: syn::TraitBoundModifiers::default(),
                        maybe: None,
                        path: bound.clone(),
                    })]))vec![syn::TypeParamBound::Trait(syn::TraitBound {
346                paren_token: None,
347                lifetimes: None,
348                modifiers: syn::TraitBoundModifiers::default(),
349                maybe: None,
350                path: bound.clone(),
351            })]
352            .into_iter()
353            .collect(),
354        }));
355    generics
356}
357
358pub fn with_lifetime_bound(generics: &syn::Generics, lifetime: &str) -> syn::Generics {
359    let bound = syn::Lifetime::new(lifetime, Span::call_site());
360    let def = syn::LifetimeParam {
361        attrs: Vec::new(),
362        lifetime: bound.clone(),
363        colon_token: None,
364        bounds: Punctuated::new(),
365    };
366
367    let params = Some(syn::GenericParam::Lifetime(def))
368        .into_iter()
369        .chain(generics.params.iter().cloned().map(|mut param| {
370            match &mut param {
371                syn::GenericParam::Lifetime(param) => {
372                    param.bounds.push(bound.clone());
373                }
374                syn::GenericParam::Type(param) => {
375                    param
376                        .bounds
377                        .push(syn::TypeParamBound::Lifetime(bound.clone()));
378                }
379                syn::GenericParam::Const(_) => {}
380            }
381            param
382        }))
383        .collect();
384
385    syn::Generics {
386        params,
387        ..generics.clone()
388    }
389}
390
391fn type_of_item(cont: &Container) -> syn::Type {
392    syn::Type::Path(syn::TypePath {
393        attrs: Vec::new(),
394        qself: None,
395        path: syn::Path {
396            leading_colon: None,
397            segments: ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [syn::PathSegment {
                    ident: cont.ident.clone(),
                    arguments: syn::PathArguments::AngleBracketed(syn::AngleBracketedGenericArguments {
                            colon2_token: None,
                            lt_token: <::syn::token::Lt>::default(),
                            args: cont.generics.params.iter().map(|param|
                                        match param {
                                            syn::GenericParam::Type(param) => {
                                                syn::GenericArgument::Type(syn::Type::Path(syn::TypePath {
                                                            attrs: Vec::new(),
                                                            qself: None,
                                                            path: param.ident.clone().into(),
                                                        }))
                                            }
                                            syn::GenericParam::Lifetime(param) => {
                                                syn::GenericArgument::Lifetime(param.lifetime.clone())
                                            }
                                            syn::GenericParam::Const(_) => {
                                                {
                                                    ::core::panicking::panic_fmt(format_args!("Serde does not support const generics yet"));
                                                };
                                            }
                                        }).collect(),
                            gt_token: <::syn::token::Gt>::default(),
                        }),
                }]))vec![syn::PathSegment {
398                ident: cont.ident.clone(),
399                arguments: syn::PathArguments::AngleBracketed(
400                    syn::AngleBracketedGenericArguments {
401                        colon2_token: None,
402                        lt_token: <Token![<]>::default(),
403                        args: cont
404                            .generics
405                            .params
406                            .iter()
407                            .map(|param| match param {
408                                syn::GenericParam::Type(param) => {
409                                    syn::GenericArgument::Type(syn::Type::Path(syn::TypePath {
410                                        attrs: Vec::new(),
411                                        qself: None,
412                                        path: param.ident.clone().into(),
413                                    }))
414                                }
415                                syn::GenericParam::Lifetime(param) => {
416                                    syn::GenericArgument::Lifetime(param.lifetime.clone())
417                                }
418                                syn::GenericParam::Const(_) => {
419                                    panic!("Serde does not support const generics yet");
420                                }
421                            })
422                            .collect(),
423                        gt_token: <Token![>]>::default(),
424                    },
425                ),
426            }]
427            .into_iter()
428            .collect(),
429        },
430    })
431}