Skip to main content

martin_config_macros/
lib.rs

1//! Derive macros for Martin's configuration types.
2
3use proc_macro::TokenStream;
4use proc_macro2::TokenStream as TokenStream2;
5use quote::quote;
6use syn::{Data, DeriveInput, Fields, GenericParam, Generics, parse_macro_input};
7
8/// Derives an empty `ConfigurationLivecycleHooks` impl, so the type opts into the trait's default hooks
9#[proc_macro_derive(ConfigurationLivecycleHooks)]
10pub fn derive_configuration_livecycle_hooks(input: TokenStream) -> TokenStream {
11    let input = parse_macro_input!(input as DeriveInput);
12    let ident = &input.ident;
13    let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
14    quote! {
15        #[automatically_derived]
16        impl #impl_generics crate::config::file::ConfigurationLivecycleHooks
17            for #ident #ty_generics #where_clause {}
18    }
19    .into()
20}
21
22/// Derives `CollectUnrecognizedKeys` for a config struct or enum.
23///
24/// Recurses into every field, except:
25/// - `#[serde(flatten)]` fields add no path segment,
26/// - `#[serde(skip)]` fields are ignored, and
27/// - `#[serde(rename)]` sets a field's path segment.
28#[proc_macro_derive(CollectUnrecognizedKeys)]
29pub fn derive_collect_unrecognized_keys(input: TokenStream) -> TokenStream {
30    let input = parse_macro_input!(input as DeriveInput);
31    expand(&input)
32        .unwrap_or_else(syn::Error::into_compile_error)
33        .into()
34}
35
36fn expand(input: &DeriveInput) -> syn::Result<TokenStream2> {
37    if let Some(attr) = container_rename_all(&input.attrs) {
38        return Err(syn::Error::new_spanned(
39            attr,
40            "CollectUnrecognizedKeys does not support `#[serde(rename_all)]` on recursed types; \
41             rename individual fields with `#[serde(rename = \"…\")]` instead",
42        ));
43    }
44
45    let body = match &input.data {
46        Data::Struct(data) => struct_body(&data.fields)?,
47        Data::Enum(data) => enum_body(data),
48        Data::Union(_) => {
49            return Err(syn::Error::new_spanned(
50                &input.ident,
51                "CollectUnrecognizedKeys cannot be derived for unions",
52            ));
53        }
54    };
55
56    let ident = &input.ident;
57    let generics = add_trait_bounds(input.generics.clone());
58    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
59
60    Ok(quote! {
61        const _: () = {
62            use crate::config::file::{CollectUnrecognizedKeys, UnrecognizedKeys};
63
64            #[automatically_derived]
65            #[allow(unused_variables)]
66            impl #impl_generics CollectUnrecognizedKeys for #ident #ty_generics #where_clause {
67                fn collect_unrecognized(&self, path: &str, out: &mut UnrecognizedKeys) {
68                    #body
69                }
70            }
71        };
72    })
73}
74
75/// Adds `T: CollectUnrecognizedKeys` to every generic type parameter.
76fn add_trait_bounds(mut generics: Generics) -> Generics {
77    for param in &mut generics.params {
78        if let GenericParam::Type(type_param) = param {
79            type_param
80                .bounds
81                .push(syn::parse_quote!(CollectUnrecognizedKeys));
82        }
83    }
84    generics
85}
86
87fn struct_body(fields: &Fields) -> syn::Result<TokenStream2> {
88    let fields = match fields {
89        Fields::Named(fields) => fields,
90        Fields::Unit => return Ok(quote! {}),
91        Fields::Unnamed(_) => {
92            return Err(syn::Error::new_spanned(
93                fields,
94                "CollectUnrecognizedKeys cannot be derived for tuple structs; implement it manually",
95            ));
96        }
97    };
98
99    let mut stmts = Vec::new();
100    for field in &fields.named {
101        if serde_flag_is_set(&field.attrs, "skip") {
102            continue;
103        }
104        let member = field.ident.as_ref().expect("named field has an ident");
105        if serde_flag_is_set(&field.attrs, "flatten") {
106            stmts.push(quote! {
107                CollectUnrecognizedKeys::collect_unrecognized(&self.#member, path, out);
108            });
109        } else {
110            let name = serde_field_name(field, member);
111            stmts.push(quote! {
112                CollectUnrecognizedKeys::collect_unrecognized(
113                    &self.#member,
114                    &format!("{path}{}.", #name),
115                    out,
116                );
117            });
118        }
119    }
120    Ok(quote! { #(#stmts)* })
121}
122
123fn enum_body(data: &syn::DataEnum) -> TokenStream2 {
124    let mut arms = Vec::new();
125    for variant in &data.variants {
126        let variant_ident = &variant.ident;
127        match &variant.fields {
128            Fields::Unit => arms.push(quote! { Self::#variant_ident => {} }),
129            Fields::Unnamed(fields) => {
130                let bindings: Vec<_> = (0..fields.unnamed.len())
131                    .map(|i| quote::format_ident!("field{i}"))
132                    .collect();
133                let recurse = bindings.iter().map(|binding| {
134                    quote! {
135                        CollectUnrecognizedKeys::collect_unrecognized(#binding, path, out);
136                    }
137                });
138                arms.push(quote! {
139                    Self::#variant_ident(#(#bindings),*) => { #(#recurse)* }
140                });
141            }
142            Fields::Named(fields) => {
143                let members: Vec<_> = fields
144                    .named
145                    .iter()
146                    .map(|f| f.ident.as_ref().expect("named field has an ident"))
147                    .collect();
148                let recurse = fields.named.iter().map(|f| {
149                    let member = f.ident.as_ref().expect("named field has an ident");
150                    let name = member.to_string();
151                    quote! {
152                        CollectUnrecognizedKeys::collect_unrecognized(
153                            #member,
154                            &format!("{path}{}.", #name),
155                            out,
156                        );
157                    }
158                });
159                arms.push(quote! {
160                    Self::#variant_ident { #(#members),* } => { #(#recurse)* }
161                });
162            }
163        }
164    }
165    quote! { match self { #(#arms)* } }
166}
167
168/// Returns `true` if any `#[serde(...)]` attribute contains the bare flag `name`.
169fn serde_flag_is_set(attrs: &[syn::Attribute], name: &str) -> bool {
170    let mut found = false;
171    for attr in attrs {
172        if !attr.path().is_ident("serde") {
173            continue;
174        }
175        let _ = attr.parse_nested_meta(|meta| {
176            if meta.path.is_ident(name) {
177                found = true;
178            }
179            if meta.input.peek(syn::Token![=]) {
180                let _: syn::Expr = meta.value()?.parse()?;
181            }
182            Ok(())
183        });
184    }
185    found
186}
187
188fn serde_field_name(field: &syn::Field, member: &syn::Ident) -> String {
189    for attr in &field.attrs {
190        if !attr.path().is_ident("serde") {
191            continue;
192        }
193        let mut rename = None;
194        let _ = attr.parse_nested_meta(|meta| {
195            if meta.path.is_ident("rename") {
196                let value = meta.value()?;
197                let lit: syn::LitStr = value.parse()?;
198                rename = Some(lit.value());
199            } else if meta.input.peek(syn::Token![=]) {
200                let _: syn::Expr = meta.value()?.parse()?;
201            }
202            Ok(())
203        });
204        if let Some(rename) = rename {
205            return rename;
206        }
207    }
208    member.to_string()
209}
210
211/// Returns the offending `#[serde(...)]` attribute if it sets `rename_all`, so the derive can reject it (unsupported).
212fn container_rename_all(attrs: &[syn::Attribute]) -> Option<&syn::Attribute> {
213    for attr in attrs {
214        if !attr.path().is_ident("serde") {
215            continue;
216        }
217        let mut found = false;
218        let _ = attr.parse_nested_meta(|meta| {
219            if meta.path.is_ident("rename_all") {
220                found = true;
221            }
222            if meta.input.peek(syn::Token![=]) {
223                let _: syn::Expr = meta.value()?.parse()?;
224            }
225            Ok(())
226        });
227        if found {
228            return Some(attr);
229        }
230    }
231    None
232}
233
234#[cfg(test)]
235mod tests {
236    use rstest::rstest;
237    use syn::parse::Parser as _;
238    use syn::{DeriveInput, Field};
239
240    use crate::{expand, serde_field_name, serde_flag_is_set};
241
242    fn parse_field(src: &str) -> Field {
243        Field::parse_named.parse_str(src).expect("field parses")
244    }
245
246    #[rstest]
247    #[case::rename_all_on_struct(
248        r#"#[serde(rename_all = "kebab-case")] struct S { a: bool }"#,
249        "does not support `#[serde(rename_all)]`"
250    )]
251    #[case::rename_all_on_enum(
252        r#"#[serde(rename_all = "kebab-case")] enum E { A }"#,
253        "does not support `#[serde(rename_all)]`"
254    )]
255    #[case::rename_all_beside_other_options(
256        r#"#[serde(deny_unknown_fields, rename_all = "kebab-case")] struct S { a: bool }"#,
257        "does not support `#[serde(rename_all)]`"
258    )]
259    #[case::rename_all_in_a_second_attribute(
260        r#"#[serde(default)] #[serde(rename_all = "kebab-case")] struct S { a: bool }"#,
261        "does not support `#[serde(rename_all)]`"
262    )]
263    #[case::union("union U { a: bool }", "cannot be derived for unions")]
264    #[case::tuple_struct("struct S(bool);", "cannot be derived for tuple structs")]
265    #[case::newtype_struct("struct S(Inner);", "cannot be derived for tuple structs")]
266    fn expand_rejects(#[case] src: &str, #[case] expected: &str) {
267        let input: DeriveInput = syn::parse_str(src).expect("input parses");
268        let err = expand(&input).expect_err("input is rejected").to_string();
269        assert!(err.contains(expected), "unexpected error: {err}");
270    }
271
272    #[rstest]
273    #[case::unit_struct("struct S;")]
274    #[case::empty_struct("struct S {}")]
275    #[case::rename_all_fields_is_a_different_option(
276        r#"#[serde(rename_all_fields = "kebab-case")] enum E { A { b: bool } }"#
277    )]
278    #[case::rename_all_on_a_variant(
279        r#"enum E { #[serde(rename_all = "kebab-case")] A { b: bool } }"#
280    )]
281    #[case::non_serde_rename_all(r#"#[schemars(rename_all = "kebab-case")] struct S { a: bool }"#)]
282    fn expand_accepts(#[case] src: &str) {
283        let input: DeriveInput = syn::parse_str(src).expect("input parses");
284        expand(&input).expect("input is accepted");
285    }
286
287    #[rstest]
288    #[case::no_attributes("a: bool", "a")]
289    #[case::rename(r#"#[serde(rename = "renamed")] a: bool"#, "renamed")]
290    #[case::rename_after_a_valued_option(
291        r#"#[serde(default = "d", rename = "renamed")] a: bool"#,
292        "renamed"
293    )]
294    #[case::rename_after_a_bare_flag(r#"#[serde(default, rename = "renamed")] a: bool"#, "renamed")]
295    #[case::rename_in_a_second_attribute(
296        r#"#[serde(default)] #[serde(rename = "renamed")] a: bool"#,
297        "renamed"
298    )]
299    #[case::rename_without_a_value("#[serde(rename)] a: bool", "a")]
300    #[case::rename_with_a_non_string_value("#[serde(rename = 7)] a: bool", "a")]
301    #[case::other_namespace(r#"#[schemars(rename = "renamed")] a: bool"#, "a")]
302    #[case::unrelated_option(r#"#[serde(alias = "renamed")] a: bool"#, "a")]
303    fn field_name(#[case] src: &str, #[case] expected: &str) {
304        let field = parse_field(src);
305        let ident = field.ident.clone().expect("field is named");
306        assert_eq!(serde_field_name(&field, &ident), expected);
307    }
308
309    #[rstest]
310    #[case::bare_flag("#[serde(flatten)] a: bool", "flatten", true)]
311    #[case::flag_after_a_valued_option(
312        r#"#[serde(default = "d", flatten)] a: bool"#,
313        "flatten",
314        true
315    )]
316    #[case::flag_before_a_valued_option(
317        r#"#[serde(flatten, default = "d")] a: bool"#,
318        "flatten",
319        true
320    )]
321    #[case::flag_in_a_second_attribute("#[serde(default)] #[serde(skip)] a: bool", "skip", true)]
322    #[case::absent("#[serde(default)] a: bool", "flatten", false)]
323    #[case::prefix_of_another_option(
324        r#"#[serde(skip_serializing_if = "f")] a: bool"#,
325        "skip",
326        false
327    )]
328    #[case::other_namespace("#[schemars(flatten)] a: bool", "flatten", false)]
329    #[case::no_attributes("a: bool", "flatten", false)]
330    fn flag_is_set(#[case] src: &str, #[case] name: &str, #[case] expected: bool) {
331        assert_eq!(serde_flag_is_set(&parse_field(src).attrs, name), expected);
332    }
333}