use heck::ToSnakeCase as _;
use quote::format_ident;
use quote::quote;
use syn::parse_quote;
use crate::ir::ConfigStructDefinition;
use crate::ir::DefaultValue;
use crate::ir::FieldDefinition;
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_default_impl<'a>(
ty: &syn::Ident,
fields: impl Iterator<Item = &'a FieldDefinition>,
) -> proc_macro2::TokenStream {
let fields = fields.flat_map(|field| {
let name = &field.name;
let default_value = field.default_value()?;
let cfg_attrs = &field.cfg_attrs;
Some(quote! {
#( #cfg_attrs )*
#name : #default_value
})
});
quote! {
#[automatically_derived]
impl ::core::default::Default for #ty {
fn default() -> Self {
Self {
#( #fields ),*
}
}
}
}
}
fn generate_validate_impl<'a>(
ty: &syn::Ident,
fields: impl Iterator<Item = &'a FieldDefinition>,
struct_validator: Option<syn::Path>,
) -> proc_macro2::TokenStream {
let validate_fields = fields.flat_map(|field| {
let name = &field.name;
let cfg_attrs = &field.cfg_attrs;
let serde_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! {
#( #cfg_attrs )*
#path(&self.#name, errors.nest(#serde_name));
}),
None => Some(quote! {
#( #cfg_attrs )*
self.#name.validate(errors.nest(#serde_name));
}),
}
});
let validate_fields = quote! {
#(#validate_fields)*
};
let body = match struct_validator {
Some(path) => quote! {
let pre_validate_len = errors.len();
#validate_fields
if errors.len() == pre_validate_len {
#path(self, errors);
}
},
None => validate_fields,
};
quote! {
#[automatically_derived]
impl ::apollo_configuration::Validate for #ty {
fn validate(&self, mut errors: ::apollo_configuration::ErrorCollector<'_>) {
#body
}
}
}
}
fn generate_field(parsed_field: FieldDefinition, default_fn: Option<&syn::Ident>) -> syn::Field {
let default_attribute = generate_serde_default_attribute(&parsed_field, default_fn);
let mut attrs = parsed_field.cfg_attrs;
attrs.extend(parsed_field.attrs);
let mut field = syn::Field {
attrs,
ident: Some(parsed_field.name.clone()),
ty: parsed_field.ty.clone(),
vis: parsed_field.vis,
mutability: syn::FieldMutability::None,
colon_token: None,
};
field.attrs.extend(default_attribute);
if !parsed_field.schemars_attrs.is_empty() {
let schemars_attrs = &parsed_field.schemars_attrs;
field.attrs.push(parse_quote! {
#[schemars( #( #schemars_attrs ),* )]
});
}
field
}
struct FieldDefaultValues {
fields: Vec<(syn::Ident, Option<syn::ItemFn>)>,
}
pub(crate) fn generate_configuration_struct(
ir: ConfigStructDefinition,
) -> proc_macro2::TokenStream {
let ConfigStructDefinition {
struct_token,
name,
vis,
generics,
attrs,
validator,
fields,
} = ir;
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
let mut field_defaults = FieldDefaultValues { fields: vec![] };
for field in &fields {
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_{name}_{field_name}");
Some(parse_quote! {
fn #ident() -> #field_ty {
#value
}
})
}
None => continue,
};
field_defaults.fields.push((field.name.clone(), default_fn));
}
let default_impl = if field_defaults.fields.len() == fields.len() {
generate_default_impl(&name, fields.iter())
} else {
Default::default()
};
let validate_impl = generate_validate_impl(&name, fields.iter(), validator);
let mut output_fields = vec![];
for def in fields {
let default_fn = field_defaults
.fields
.iter()
.find(|(field_name, _)| *field_name == def.name)
.and_then(|(_, default_fn)| default_fn.as_ref())
.map(|default_fn| &default_fn.sig.ident);
output_fields.push(generate_field(def, default_fn));
}
let output_fields = syn::Fields::Named(syn::FieldsNamed {
brace_token: Default::default(),
named: output_fields.into_iter().collect(),
});
let auxiliary_definitions = field_defaults
.fields
.into_iter()
.flat_map(|(_, default_fn)| default_fn);
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(deny_unknown_fields)]
#vis #struct_token #name #generics
#output_fields
#( #auxiliary_definitions )*
#default_impl
#validate_impl
impl #impl_generics ::apollo_configuration::Configuration for #name #ty_generics #where_clause {}
}
}