use proc_macro::TokenStream;
use quote::{format_ident, quote};
use syn::punctuated::Punctuated;
use syn::{
Attribute, Error, Expr, Field, Fields, ItemStruct, LitStr, Meta, Token, parse_macro_input,
parse_quote,
};
#[proc_macro_attribute]
pub fn serde_const_default(attr: TokenStream, item: TokenStream) -> TokenStream {
if !attr.is_empty() {
return Error::new_spanned(
proc_macro2::TokenStream::from(attr),
"`serde_const_default` does not accept arguments",
)
.to_compile_error()
.into();
}
let item = parse_macro_input!(item as ItemStruct);
expand_serde_const_default(item)
.unwrap_or_else(Error::into_compile_error)
.into()
}
fn expand_serde_const_default(mut item: ItemStruct) -> Result<proc_macro2::TokenStream, Error> {
let mut helpers = Vec::new();
let struct_ident = item.ident.clone();
match &mut item.fields {
Fields::Named(fields) => {
for (index, field) in fields.named.iter_mut().enumerate() {
if let Some(helper) = rewrite_field(&struct_ident, index, field)? {
helpers.push(helper);
}
}
}
Fields::Unnamed(fields) => {
for field in &fields.unnamed {
if has_const_default_attr(field) {
return Err(Error::new_spanned(
field,
"`serde_const_default` only supports named fields",
));
}
}
}
Fields::Unit => {}
}
Ok(quote! {
#(#helpers)*
#item
})
}
fn rewrite_field(
struct_ident: &syn::Ident,
index: usize,
field: &mut Field,
) -> Result<Option<proc_macro2::TokenStream>, Error> {
let Some(default) = take_const_default_attr(&mut field.attrs)? else {
return Ok(None);
};
if let Some(attr) = field.attrs.iter().find(|attr| serde_attr_has_default(attr)) {
return Err(Error::new_spanned(
attr,
"`const_default` cannot be combined with `#[serde(default)]`",
));
}
let Some(field_ident) = field.ident.as_ref() else {
return Err(Error::new_spanned(
&*field,
"`serde_const_default` only supports named fields",
));
};
let helper_ident = format_ident!(
"__serde_const_default_{}_{}_{}",
struct_ident,
field_ident,
index
);
let helper_name = LitStr::new(&helper_ident.to_string(), helper_ident.span());
let ty = &field.ty;
let expr = default.expr;
let (function_qualifier, body) = match default.kind {
DefaultKind::Plain => (quote! { const }, quote! { #expr }),
DefaultKind::From => (quote! {}, quote! { ::core::convert::From::from(#expr) }),
};
field.attrs.push(parse_quote! {
#[serde(default = #helper_name)]
});
Ok(Some(quote! {
#[allow(non_snake_case)]
#function_qualifier fn #helper_ident() -> #ty {
#body
}
}))
}
fn take_const_default_attr(attrs: &mut Vec<Attribute>) -> Result<Option<ConstDefault>, Error> {
let mut default = None;
let mut error = None::<Error>;
let mut retained = Vec::with_capacity(attrs.len());
for attr in attrs.drain(..) {
match parse_const_default_attr(&attr)? {
Some(parsed) => {
if default.is_some() {
let duplicate_error = Error::new_spanned(
attr,
"only one const default attribute is allowed per field",
);
if let Some(error) = &mut error {
error.combine(duplicate_error);
} else {
error = Some(duplicate_error);
}
} else {
default = Some(parsed);
}
}
None => retained.push(attr),
}
}
*attrs = retained;
if let Some(error) = error {
Err(error)
} else {
Ok(default)
}
}
fn parse_const_default_attr(attr: &Attribute) -> Result<Option<ConstDefault>, Error> {
if attr.path().is_ident("const_default") {
let Meta::NameValue(name_value) = &attr.meta else {
return Err(Error::new_spanned(
attr,
"expected `#[const_default = EXPR]`",
));
};
return Ok(Some(ConstDefault {
kind: DefaultKind::Plain,
expr: name_value.value.clone(),
}));
}
if attr.path().is_ident("const_default_from") {
return Ok(Some(ConstDefault {
kind: DefaultKind::From,
expr: attr.parse_args::<Expr>()?,
}));
}
Ok(None)
}
fn has_const_default_attr(field: &Field) -> bool {
field.attrs.iter().any(|attr| {
attr.path().is_ident("const_default") || attr.path().is_ident("const_default_from")
})
}
fn serde_attr_has_default(attr: &Attribute) -> bool {
if !attr.path().is_ident("serde") {
return false;
}
let Meta::List(list) = &attr.meta else {
return false;
};
list.parse_args_with(Punctuated::<Meta, Token![,]>::parse_terminated)
.map(|metas| {
metas.iter().any(|meta| match meta {
Meta::Path(path) => path.is_ident("default"),
Meta::NameValue(name_value) => name_value.path.is_ident("default"),
Meta::List(list) => list.path.is_ident("default"),
})
})
.unwrap_or(false)
}
struct ConstDefault {
kind: DefaultKind,
expr: Expr,
}
enum DefaultKind {
Plain,
From,
}