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::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);

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

    // Combine cfg_attrs with other attrs - cfg_attrs go first so they apply to the 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![];

            // Generate auxiliary default functions for fields with custom defaults
            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));
            }

            // Generate output fields
            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();

            // Collect auxiliary functions
            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;

    // Generate the default value based on the variant's fields
    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;
        // 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_variant_name = variant.name.to_string().to_snake_case();

        match &variant.fields {
            VariantFields::Unit => quote::quote! {
                #( #cfg_attrs )*
                Self::#name => {}
            },
            // Single-element tuple means there's no nesting in the YAML representation
            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) => {
                // XXX(@goto-bus-stop): This is currently _always assuming_ the default serde
                // strategy of serializing enums as objects with a single property (the variant
                // name). It will produce incorrect code for other serialization strategies. We
                // will want to support other (untagged, tagged-by-property) strategies in the
                // future.
                let validators = fields.iter().flat_map(|field| {
                    let field_name = &field.name;
                    let field_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_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 {
        // With `#[configuration(validate = validate_fn)]`, `validate_fn` is called on the whole
        // enum after all fields are validated.
        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();

    // Find the default variant, if any
    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);

    // Generate variants and collect auxiliary functions
    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 {}
    }
}