1use proc_macro::TokenStream;
4use proc_macro2::TokenStream as TokenStream2;
5use quote::quote;
6use syn::{Data, DeriveInput, Fields, GenericParam, Generics, parse_macro_input};
7
8#[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#[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
75fn 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
168fn 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
211fn 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}