Skip to content

Commit f625043

Browse files
Merge pull request #81 from balena-io-modules/mahler-derive-default
Add `default` attribute to State derive macro
2 parents 6260442 + 2742048 commit f625043

2 files changed

Lines changed: 177 additions & 129 deletions

File tree

mahler-derive/src/lib.rs

Lines changed: 99 additions & 129 deletions
Original file line numberDiff line numberDiff line change
@@ -31,62 +31,62 @@
3131
3232
use proc_macro::TokenStream;
3333
use quote::{format_ident, quote};
34-
use syn::{parse_macro_input, Attribute, Data, DeriveInput, Error, Field, Fields, Meta, Result};
35-
36-
type NamedTargetFields = Vec<proc_macro2::TokenStream>;
37-
type UnnamedTargetFields = Vec<proc_macro2::TokenStream>;
34+
use syn::{
35+
parse::Parser, parse_macro_input, punctuated::Punctuated, Attribute, Data, DeriveInput, Error,
36+
Field, Fields, Ident, Meta, Result, Token,
37+
};
3838

3939
// ============================================================================
4040
// Parsing functions
4141
// ============================================================================
4242

43-
/// Parse whether a field has the #[mahler(internal)] attribute
44-
fn has_mahler_internal_attribute(attrs: &[Attribute]) -> Result<bool> {
43+
/// Parsed mahler field attributes
44+
struct MahlerFieldAttrs {
45+
internal: bool,
46+
default: bool,
47+
}
48+
49+
/// Parse mahler field attributes, returning which flags are present
50+
fn parse_mahler_field_attrs(attrs: &[Attribute]) -> Result<MahlerFieldAttrs> {
51+
let mut result = MahlerFieldAttrs {
52+
internal: false,
53+
default: false,
54+
};
4555
for attr in attrs {
4656
if attr.path().is_ident("mahler") {
47-
match &attr.meta {
48-
Meta::List(meta_list) => {
49-
let tokens = &meta_list.tokens;
50-
if tokens.to_string().trim() == "internal" {
51-
return Ok(true);
57+
if let Meta::List(meta_list) = &attr.meta {
58+
let parser = Punctuated::<Ident, Token![,]>::parse_terminated;
59+
let idents = parser.parse2(meta_list.tokens.clone())?;
60+
for ident in &idents {
61+
if ident == "internal" {
62+
result.internal = true;
63+
} else if ident == "default" {
64+
result.default = true;
65+
} else {
66+
return Err(Error::new_spanned(
67+
ident,
68+
format!("unknown mahler field attribute: `{ident}`"),
69+
));
5270
}
5371
}
54-
Meta::Path(_) => {}
55-
Meta::NameValue(_) => {}
5672
}
5773
}
5874
}
59-
Ok(false)
75+
Ok(result)
6076
}
6177

6278
/// Extract additional derives from #[mahler(derive(...))] attribute
6379
fn extract_mahler_derives(attrs: &[Attribute]) -> Result<Option<proc_macro2::TokenStream>> {
6480
for attr in attrs {
6581
if attr.path().is_ident("mahler") {
6682
if let Meta::List(meta_list) = &attr.meta {
67-
let tokens_str = meta_list.tokens.to_string();
68-
69-
if let Some(derives_part) = tokens_str.strip_prefix("derive") {
70-
let derives_part = derives_part.trim();
71-
if derives_part.starts_with('(') && derives_part.ends_with(')') {
72-
let inner = &derives_part[1..derives_part.len() - 1];
73-
74-
let traits: Vec<_> = inner
75-
.split(',')
76-
.map(|s| s.trim())
77-
.filter(|s| !s.is_empty())
78-
.collect();
79-
83+
let inner_meta: Meta = syn::parse2(meta_list.tokens.clone())?;
84+
if let Meta::List(inner_list) = inner_meta {
85+
if inner_list.path.is_ident("derive") {
86+
let parser = Punctuated::<Ident, Token![,]>::parse_terminated;
87+
let traits = parser.parse2(inner_list.tokens.clone())?;
8088
if !traits.is_empty() {
81-
let trait_tokens: Vec<proc_macro2::TokenStream> = traits
82-
.iter()
83-
.map(|t| {
84-
let ident = syn::Ident::new(t, proc_macro2::Span::call_site());
85-
quote! { #ident }
86-
})
87-
.collect();
88-
89-
return Ok(Some(quote! { #(#trait_tokens),* }));
89+
return Ok(Some(quote! { #traits }));
9090
}
9191
}
9292
}
@@ -130,18 +130,18 @@ fn validate_enum_variants(data_enum: &syn::DataEnum) -> Result<()> {
130130

131131
/// Process named fields to generate target field definitions
132132
fn process_named_fields(
133-
fields: &syn::punctuated::Punctuated<Field, syn::Token![,]>,
134-
) -> Result<NamedTargetFields> {
135-
let mut target_fields = NamedTargetFields::new();
133+
fields: &Punctuated<Field, Token![,]>,
134+
) -> Result<Vec<proc_macro2::TokenStream>> {
135+
let mut target_fields = Vec::new();
136136

137137
for field in fields {
138-
let field_name = field.ident.as_ref().unwrap();
138+
let field_name = field.ident.as_ref().expect("named field has ident");
139139
let field_type = &field.ty;
140140
let field_vis = &field.vis;
141141

142-
let is_internal = has_mahler_internal_attribute(&field.attrs)?;
142+
let attrs = parse_mahler_field_attrs(&field.attrs)?;
143143

144-
if !is_internal {
144+
if !attrs.internal {
145145
let target_attrs = filter_field_attributes(&field.attrs);
146146

147147
target_fields.push(quote! {
@@ -156,16 +156,16 @@ fn process_named_fields(
156156

157157
/// Process unnamed fields to generate target field definitions
158158
fn process_unnamed_fields(
159-
fields: &syn::punctuated::Punctuated<Field, syn::Token![,]>,
160-
) -> Result<UnnamedTargetFields> {
161-
let mut target_fields = UnnamedTargetFields::new();
159+
fields: &Punctuated<Field, Token![,]>,
160+
) -> Result<Vec<proc_macro2::TokenStream>> {
161+
let mut target_fields = Vec::new();
162162

163163
for field in fields.iter() {
164164
let field_type = &field.ty;
165165

166-
let is_internal = has_mahler_internal_attribute(&field.attrs)?;
166+
let attrs = parse_mahler_field_attrs(&field.attrs)?;
167167

168-
if is_internal {
168+
if attrs.internal {
169169
return Err(Error::new_spanned(
170170
field,
171171
"#[mahler(internal)] is only supported on named struct fields, not tuple struct fields",
@@ -204,6 +204,41 @@ fn is_collection_type(ty: &syn::Type) -> bool {
204204
false
205205
}
206206

207+
/// Build `impl<'de, ...>` generics and where clause for Deserialize impls
208+
fn de_generics(generics: &syn::Generics) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
209+
let (_, _, where_clause) = generics.split_for_impl();
210+
let de_impl_generics = if generics.params.is_empty() {
211+
quote! { <'de> }
212+
} else {
213+
let params = &generics.params;
214+
quote! { <'de, #params> }
215+
};
216+
let de_where_clause = if let Some(wc) = where_clause {
217+
quote! { #wc }
218+
} else {
219+
quote! {}
220+
};
221+
(de_impl_generics, de_where_clause)
222+
}
223+
224+
/// Filter fields, optionally excluding those marked `#[mahler(internal)]`
225+
fn relevant_fields<'a>(
226+
fields: impl Iterator<Item = &'a Field>,
227+
exclude_internal: bool,
228+
) -> Vec<&'a Field> {
229+
if exclude_internal {
230+
fields
231+
.filter(|f| {
232+
!parse_mahler_field_attrs(&f.attrs)
233+
.map(|a| a.internal)
234+
.unwrap_or(false)
235+
})
236+
.collect()
237+
} else {
238+
fields.collect()
239+
}
240+
}
241+
207242
// ============================================================================
208243
// Code generation functions
209244
// ============================================================================
@@ -216,22 +251,12 @@ fn generate_serialize_impl(
216251
exclude_internal: bool,
217252
) -> proc_macro2::TokenStream {
218253
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
219-
let all_fields: Vec<_> = fields.iter().collect();
220-
221-
let relevant_fields: Vec<_> = if exclude_internal {
222-
all_fields
223-
.iter()
224-
.filter(|f| !has_mahler_internal_attribute(&f.attrs).unwrap_or(false))
225-
.copied()
226-
.collect()
227-
} else {
228-
all_fields
229-
};
254+
let relevant = relevant_fields(fields.iter(), exclude_internal);
230255

231-
let field_count = relevant_fields.len();
256+
let field_count = relevant.len();
232257
let struct_name_str = struct_name.to_string();
233258

234-
let field_serializations: Vec<_> = relevant_fields
259+
let field_serializations: Vec<_> = relevant
235260
.iter()
236261
.filter_map(|field| {
237262
let field_name = field.ident.as_ref()?;
@@ -276,33 +301,21 @@ fn generate_deserialize_impl(
276301
exclude_internal: bool,
277302
) -> proc_macro2::TokenStream {
278303
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
279-
let all_fields: Vec<_> = fields.iter().collect();
304+
let relevant = relevant_fields(fields.iter(), exclude_internal);
280305

281-
let relevant_fields: Vec<_> = if exclude_internal {
282-
all_fields
283-
.iter()
284-
.filter(|f| !has_mahler_internal_attribute(&f.attrs).unwrap_or(false))
285-
.copied()
286-
.collect()
287-
} else {
288-
all_fields
289-
};
290-
291-
let field_names: Vec<_> = relevant_fields
292-
.iter()
293-
.filter_map(|f| f.ident.as_ref())
294-
.collect();
306+
let field_names: Vec<_> = relevant.iter().filter_map(|f| f.ident.as_ref()).collect();
295307
let field_strs: Vec<_> = field_names.iter().map(|n| n.to_string()).collect();
296308
let struct_name_str = struct_name.to_string();
297309

298-
// Generate field assignments based on whether they're Option or collection types
299-
let field_assignments: Vec<_> = relevant_fields
310+
// Generate field assignments based on whether they're Option, collection, or #[mahler(default)] types
311+
let field_assignments: Vec<_> = relevant
300312
.iter()
301313
.filter_map(|field| {
302314
let field_name = field.ident.as_ref()?;
303315
let field_str = field_name.to_string();
316+
let attrs = parse_mahler_field_attrs(&field.attrs).ok()?;
304317

305-
if is_option_type(&field.ty) || is_collection_type(&field.ty) {
318+
if is_option_type(&field.ty) || is_collection_type(&field.ty) || attrs.default {
306319
Some(quote! {
307320
#field_name: #field_name.unwrap_or_default()
308321
})
@@ -314,19 +327,7 @@ fn generate_deserialize_impl(
314327
})
315328
.collect();
316329

317-
// Handle generics with the de lifetime
318-
let de_impl_generics = if generics.params.is_empty() {
319-
quote! { <'de> }
320-
} else {
321-
let params = &generics.params;
322-
quote! { <'de, #params> }
323-
};
324-
325-
let de_where_clause = if let Some(wc) = where_clause {
326-
quote! { #wc }
327-
} else {
328-
quote! {}
329-
};
330+
let (de_impl_generics, de_where_clause) = de_generics(generics);
330331

331332
quote! {
332333
impl #de_impl_generics ::mahler::serde::Deserialize<'de> for #struct_name #ty_generics #de_where_clause {
@@ -494,18 +495,7 @@ fn expand_state_derive(input: DeriveInput) -> Result<TokenStream> {
494495
let target_fields = process_unnamed_fields(&fields.unnamed)?;
495496
let struct_attrs = filter_attributes(&input.attrs);
496497
let field_type = &fields.unnamed.iter().next().unwrap().ty;
497-
let de_impl_generics = if generics.params.is_empty() {
498-
quote! { <'de> }
499-
} else {
500-
let params = &generics.params;
501-
quote! { <'de, #params> }
502-
};
503-
504-
let de_where_clause = if let Some(wc) = where_clause {
505-
quote! { #wc }
506-
} else {
507-
quote! {}
508-
};
498+
let (de_impl_generics, de_where_clause) = de_generics(generics);
509499

510500
let expanded = quote! {
511501
impl #impl_generics ::mahler::serde::Serialize for #struct_name #ty_generics #where_clause {
@@ -564,19 +554,7 @@ fn expand_state_derive(input: DeriveInput) -> Result<TokenStream> {
564554
}
565555

566556
let struct_name_str = struct_name.to_string();
567-
568-
let de_impl_generics = if generics.params.is_empty() {
569-
quote! { <'de> }
570-
} else {
571-
let params = &generics.params;
572-
quote! { <'de, #params> }
573-
};
574-
575-
let de_where_clause = if let Some(wc) = where_clause {
576-
quote! { #wc }
577-
} else {
578-
quote! {}
579-
};
557+
let (de_impl_generics, de_where_clause) = de_generics(generics);
580558

581559
let expanded = quote! {
582560
impl #impl_generics ::mahler::serde::Serialize for #struct_name #ty_generics #where_clause {
@@ -639,19 +617,7 @@ fn expand_state_derive(input: DeriveInput) -> Result<TokenStream> {
639617
let variants: Vec<_> = data_enum.variants.iter().collect();
640618
let variant_names: Vec<_> = variants.iter().map(|v| &v.ident).collect();
641619
let variant_strs: Vec<_> = variant_names.iter().map(|n| n.to_string()).collect();
642-
643-
let de_impl_generics = if generics.params.is_empty() {
644-
quote! { <'de> }
645-
} else {
646-
let params = &generics.params;
647-
quote! { <'de, #params> }
648-
};
649-
650-
let de_where_clause = if let Some(wc) = where_clause {
651-
quote! { #wc }
652-
} else {
653-
quote! {}
654-
};
620+
let (de_impl_generics, de_where_clause) = de_generics(generics);
655621

656622
let expanded = quote! {
657623
impl #impl_generics ::mahler::serde::Serialize for #struct_name #ty_generics #where_clause {
@@ -801,6 +767,10 @@ fn expand_state_derive(input: DeriveInput) -> Result<TokenStream> {
801767
///
802768
/// - `#[mahler(internal)]` - Marks a struct field as internal-only, excluding it from the target type.
803769
/// Only supported on named struct fields, not tuple struct, unit struct, or enum fields.
770+
/// - `#[mahler(default)]` - Uses `Default::default()` when the field is missing during deserialization,
771+
/// instead of returning an error. This attribute only affects deserialization; during serialization,
772+
/// the field is always emitted. The field's type must implement `Default`. Can be combined
773+
/// with `internal`: `#[mahler(internal, default)]`.
804774
/// - `#[mahler(derive(Trait1, Trait2, ...))]` - Adds additional derives to the generated target struct.
805775
/// Only applies when a new target struct is created (i.e. when the source structure is not an
806776
/// enum or a unit type).

0 commit comments

Comments
 (0)