apollo-configuration-macros 0.1.1

Supporting macros for apollo-configuration (internal, do not use directly)
Documentation
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;
        // XXX(@goto-bus-stop) Using `.to_snake_case()` because that's what we default _all_ fields to,
        // but this is not 100% accurate: a user could `#[serde(rename)]` individual fields and
        // then this is wrong.
        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 {
        // With `#[configuration(validate = validate_fn)]`, `validate_fn` is called on the whole
        // structure after all fields are validated.
        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
            }
        }
    }
}

/// Returns a field definition with serde attributes applied.
fn generate_field(parsed_field: FieldDefinition, default_fn: Option<&syn::Ident>) -> syn::Field {
    let default_attribute = generate_serde_default_attribute(&parsed_field, default_fn);

    // Combine cfg_attrs with other attrs - cfg_attrs go first so they apply to the field
    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));
    }

    // Only generate a `Default` impl if there are no required fields in the type
    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 {}
    }
}