use heck::ToSnakeCase as _;
use quote::format_ident;
use quote::quote;
use syn::parse_quote;
use crate::ir::ConfigEnumDefinition;
use crate::ir::DefaultValue;
use crate::ir::FieldDefinition;
use crate::ir::VariantDefinition;
use crate::ir::VariantFields;
fn generate_serde_default_attribute(
field: &FieldDefinition,
default_fn: Option<&syn::Ident>,
) -> Option<syn::Attribute> {
match field.default_value()? {
DefaultValue::Default => Some(parse_quote! {
#[serde(default)]
}),
DefaultValue::Value(_value) => {
let ident = default_fn.expect("custom value means an auxiliary function was defined");
let ident_ref = format!("{ident}");
Some(parse_quote! { #[serde(default = #ident_ref)] })
}
}
}
fn generate_field(field: FieldDefinition, default_fn: Option<&syn::Ident>) -> syn::Field {
let default_attribute = generate_serde_default_attribute(&field, default_fn);
let mut attrs = field.cfg_attrs;
attrs.extend(field.attrs);
let mut output = syn::Field {
attrs,
ident: Some(field.name),
ty: field.ty,
vis: field.vis,
mutability: syn::FieldMutability::None,
colon_token: None,
};
output.attrs.extend(default_attribute);
if !field.schemars_attrs.is_empty() {
let schemars_attrs = &field.schemars_attrs;
output.attrs.push(parse_quote! {
#[schemars( #( #schemars_attrs ),* )]
});
}
output
}
struct VariantOutput {
variant: syn::Variant,
auxiliary_fns: Vec<syn::ItemFn>,
}
fn generate_variant(enum_name: &syn::Ident, variant: VariantDefinition) -> VariantOutput {
let VariantDefinition {
name: variant_name,
attrs,
cfg_attrs,
fields,
is_default: _,
} = variant;
let combined_attrs = {
let mut combined = cfg_attrs;
combined.extend(attrs);
combined
};
match fields {
VariantFields::Unit => VariantOutput {
variant: syn::Variant {
attrs: combined_attrs,
ident: variant_name,
fields: syn::Fields::Unit,
discriminant: None,
},
auxiliary_fns: vec![],
},
VariantFields::Tuple(fields) => VariantOutput {
variant: syn::Variant {
attrs: combined_attrs,
ident: variant_name,
fields: syn::Fields::Unnamed(fields),
discriminant: None,
},
auxiliary_fns: vec![],
},
VariantFields::Struct(field_defs) => {
let mut auxiliary_fns = vec![];
let mut field_default_fns: Vec<(syn::Ident, Option<syn::ItemFn>)> = vec![];
for field in &field_defs {
let field_ty = &field.ty;
let field_name = &field.name;
let default_fn = match field.default_value() {
Some(DefaultValue::Default) => None,
Some(DefaultValue::Value(value)) => {
let ident = format_ident!(
"_apollo_configuration_default_value_{enum_name}_{variant_name}_{field_name}"
);
Some(parse_quote! {
fn #ident() -> #field_ty {
#value
}
})
}
None => continue,
};
field_default_fns.push((field.name.clone(), default_fn));
}
let output_fields: Vec<_> = field_defs
.into_iter()
.map(|field| {
let default_fn = field_default_fns
.iter()
.find(|(name, _)| *name == field.name)
.and_then(|(_, f)| f.as_ref())
.map(|f| &f.sig.ident);
generate_field(field, default_fn)
})
.collect();
for (_, default_fn) in field_default_fns {
if let Some(f) = default_fn {
auxiliary_fns.push(f);
}
}
VariantOutput {
variant: syn::Variant {
attrs: combined_attrs,
ident: variant_name,
fields: syn::Fields::Named(syn::FieldsNamed {
brace_token: Default::default(),
named: output_fields.into_iter().collect(),
}),
discriminant: None,
},
auxiliary_fns,
}
}
}
}
fn generate_default_impl(
name: &syn::Ident,
generics: &syn::Generics,
default_variant: &VariantDefinition,
) -> proc_macro2::TokenStream {
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
let variant_name = &default_variant.name;
let default_value = match &default_variant.fields {
VariantFields::Unit => quote! { Self::#variant_name },
VariantFields::Tuple(_) => quote! { Self::#variant_name(Default::default()) },
VariantFields::Struct(fields) => {
let field_defaults = fields.iter().map(|f| {
let field_name = &f.name;
let cfg_attrs = &f.cfg_attrs;
let value = f
.default_value()
.expect("default variant struct fields must have defaults");
quote! {
#( #cfg_attrs )*
#field_name: #value
}
});
quote! { Self::#variant_name { #( #field_defaults ),* } }
}
};
quote! {
#[automatically_derived]
impl #impl_generics ::core::default::Default for #name #ty_generics #where_clause {
fn default() -> Self {
#default_value
}
}
}
}
fn generate_validate_impl<'a>(
ty: &syn::Ident,
variants: impl Iterator<Item = &'a VariantDefinition>,
enum_validator: Option<syn::Path>,
) -> proc_macro2::TokenStream {
let branches = variants.map(|variant| {
let name = &variant.name;
let cfg_attrs = &variant.cfg_attrs;
let serde_variant_name = variant.name.to_string().to_snake_case();
match &variant.fields {
VariantFields::Unit => quote::quote! {
#( #cfg_attrs )*
Self::#name => {}
},
VariantFields::Tuple(unnamed) if unnamed.unnamed.len() == 1 => {
quote::quote! {
#( #cfg_attrs )*
Self::#name ( inner ) => inner.validate(errors.nest(#serde_variant_name)),
}
}
VariantFields::Tuple(unnamed) => {
let (destructure, validators): (Vec<_>, Vec<_>) = unnamed
.unnamed
.iter()
.enumerate()
.map(|(index, _field)| {
let ident = format_ident!("f{index}");
(
ident.clone(),
quote::quote! {
#ident.validate(errors.nest(#index));
},
)
})
.unzip();
quote::quote! {
#( #cfg_attrs )*
Self::#name ( #( #destructure ),* ) => {
let mut errors = errors.nest(#serde_variant_name);
#( #validators )*
}
}
}
VariantFields::Struct(fields) => {
let validators = fields.iter().flat_map(|field| {
let field_name = &field.name;
let field_cfg_attrs = &field.cfg_attrs;
let serde_field_name = field.name.to_string().to_snake_case();
match &field.validator {
Some(crate::ir::FieldValidator::Skip) => None,
Some(crate::ir::FieldValidator::Custom(path)) => Some(quote! {
#( #field_cfg_attrs )*
#path(#field_name, errors.nest(#serde_field_name));
}),
None => Some(quote! {
#( #field_cfg_attrs )*
#field_name.validate(errors.nest(#serde_field_name));
}),
}
});
let destructure = fields.iter().map(|field| {
let field_name = &field.name;
let field_cfg_attrs = &field.cfg_attrs;
quote! { #( #field_cfg_attrs )* #field_name }
});
quote::quote! {
#( #cfg_attrs )*
Self::#name { #( #destructure ),* } => {
let mut errors = errors.nest(#serde_variant_name);
#( #validators )*
}
}
}
}
});
let validate_branches = quote::quote! {
match self {
#(#branches)*
}
};
let body = match enum_validator {
Some(path) => quote! {
let pre_validate_len = errors.len();
#validate_branches
if errors.len() == pre_validate_len {
#path(self, errors);
}
},
None => validate_branches,
};
quote! {
#[automatically_derived]
impl ::apollo_configuration::Validate for #ty {
fn validate(&self, mut errors: ::apollo_configuration::ErrorCollector<'_>) {
#body
}
}
}
}
pub(crate) fn generate_configuration_enum(ir: ConfigEnumDefinition) -> proc_macro2::TokenStream {
let ConfigEnumDefinition {
enum_token,
name,
vis,
generics,
attrs,
derive_helper_attrs,
variants,
validator,
} = ir;
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
let default_variant = variants.iter().find(|v| v.is_default);
let default_impl = default_variant
.map(|v| generate_default_impl(&name, &generics, v))
.unwrap_or_default();
let validate_impl = generate_validate_impl(&name, variants.iter(), validator);
let mut output_variants = vec![];
let mut all_auxiliary_fns = vec![];
for variant in variants {
let output = generate_variant(&name, variant);
output_variants.push(output.variant);
all_auxiliary_fns.extend(output.auxiliary_fns);
}
quote! {
#(#attrs)*
#[derive(
::std::fmt::Debug,
::std::clone::Clone,
::apollo_configuration::private::schemars::JsonSchema,
::apollo_configuration::private::serde::Deserialize,
)]
#[serde(crate = "::apollo_configuration::private::serde")]
#[schemars(crate = "::apollo_configuration::private::schemars")]
#[serde(rename_all = "snake_case", deny_unknown_fields)]
#(#derive_helper_attrs)*
#vis #enum_token #name #generics {
#( #output_variants ),*
}
#( #all_auxiliary_fns )*
#default_impl
#validate_impl
impl #impl_generics ::apollo_configuration::Configuration for #name #ty_generics #where_clause {}
}
}