inspect-derive 0.1.0

Derive macro for the inspect-rs introspection system
Documentation
//! Derive macro for the `Inspect` trait.

use proc_macro::TokenStream;
use quote::quote;
use syn::{Data, DeriveInput, Fields, parse_macro_input};

#[proc_macro_derive(Inspect, attributes(inspect))]
pub fn derive_inspect(input: TokenStream) -> TokenStream {
    let input = parse_macro_input!(input as DeriveInput);

    let name = &input.ident;
    let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();

    // Build where clause with Inspect bounds
    let mut where_clause = where_clause.cloned().unwrap_or_else(|| syn::parse_quote!(where));

    for param in &input.generics.params {
        if let syn::GenericParam::Type(type_param) = param {
            let ident = &type_param.ident;
            where_clause.predicates.push(syn::parse_quote!(#ident: inspect_core::Inspect));
        }
    }

    let inspect_impl = match &input.data {
        Data::Struct(data_struct) => impl_struct(name, &data_struct.fields),
        Data::Enum(data_enum) => impl_enum(name, data_enum),
        Data::Union(_) => {
            return syn::Error::new_spanned(name, "Inspect cannot be derived for unions")
                .to_compile_error()
                .into();
        }
    };

    let expanded = quote! {
        impl #impl_generics inspect_core::Inspect for #name #ty_generics #where_clause {
            fn inspect(&self, cx: &mut inspect_core::InspectCx<'_>) -> inspect_core::ValueRef<'_> {
                #inspect_impl
            }
        }
    };

    TokenStream::from(expanded)
}

fn impl_struct(name: &syn::Ident, fields: &Fields) -> proc_macro2::TokenStream {
    let type_name = name.to_string();

    match fields {
        Fields::Named(fields_named) => {
            let field_inspections: Vec<_> = fields_named
                .named
                .iter()
                .enumerate()
                .filter_map(|(idx, field)| {
                    let attrs = FieldAttributes::parse(&field.attrs);

                    if attrs.skip {
                        return None;
                    }

                    let field_name = field.ident.as_ref().unwrap();
                    let field_name_str = attrs.rename.unwrap_or_else(|| field_name.to_string());

                    let sensitivity = match (attrs.secret, attrs.sensitive) {
                        (true, _) => quote!(inspect_core::Sensitivity::Secret),
                        (_, true) => quote!(inspect_core::Sensitivity::Sensitive),
                        _ => quote!(inspect_core::Sensitivity::Normal),
                    };

                    Some(quote! {
                        (
                            inspect_core::FieldInfo::named(#field_name_str, #idx)
                                .with_sensitivity(#sensitivity),
                            self.#field_name.inspect(cx)
                        )
                    })
                })
                .collect();

            quote! {
                cx.visit_node();
                let fields = vec![#(#field_inspections),*];

                inspect_core::ValueRef::with_children(
                    inspect_core::Kind::Struct,
                    inspect_core::TypeInfo::new(#type_name),
                    inspect_core::Children::direct(fields),
                )
            }
        }
        Fields::Unnamed(fields_unnamed) => {
            let field_count = fields_unnamed.unnamed.len();
            let field_inspections: Vec<_> = (0..field_count)
                .map(|idx| {
                    let idx_token = syn::Index::from(idx);
                    quote! {
                        (
                            inspect_core::FieldInfo::tuple(#idx),
                            self.#idx_token.inspect(cx)
                        )
                    }
                })
                .collect();

            quote! {
                cx.visit_node();
                let fields = vec![#(#field_inspections),*];

                inspect_core::ValueRef::with_children(
                    inspect_core::Kind::TupleStruct,
                    inspect_core::TypeInfo::new(#type_name),
                    inspect_core::Children::direct(fields),
                )
            }
        }
        Fields::Unit => {
            quote! {
                cx.visit_node();
                inspect_core::ValueRef::with_type(
                    inspect_core::Kind::Struct,
                    inspect_core::TypeInfo::new(#type_name),
                )
            }
        }
    }
}

fn impl_enum(name: &syn::Ident, data_enum: &syn::DataEnum) -> proc_macro2::TokenStream {
    let type_name = name.to_string();

    let variant_arms: Vec<_> = data_enum
        .variants
        .iter()
        .enumerate()
        .map(|(variant_idx, variant)| {
            let variant_name = &variant.ident;
            let variant_name_str = variant_name.to_string();

            match &variant.fields {
                Fields::Named(fields) => {
                    let field_names: Vec<_> = fields.named.iter()
                        .map(|f| f.ident.as_ref().unwrap())
                        .collect();

                    let field_inspections: Vec<_> = fields.named.iter()
                        .enumerate()
                        .map(|(idx, field)| {
                            let field_name = field.ident.as_ref().unwrap();
                            let field_name_str = field_name.to_string();

                            quote! {
                                (
                                    inspect_core::FieldInfo::named(#field_name_str, #idx),
                                    #field_name.inspect(cx)
                                )
                            }
                        })
                        .collect();

                    quote! {
                        #name::#variant_name { #(#field_names),* } => {
                            cx.visit_node();
                            let variant = inspect_core::VariantInfo::new(#variant_name_str, #variant_idx);
                            let fields = vec![#(#field_inspections),*];

                            inspect_core::ValueRef::with_children(
                                inspect_core::Kind::Enum,
                                inspect_core::TypeInfo::new(#type_name),
                                inspect_core::Children::direct(fields),
                            )
                            .with_variant(variant)
                        }
                    }
                }
                Fields::Unnamed(fields) => {
                    let field_count = fields.unnamed.len();
                    let field_bindings: Vec<_> = (0..field_count)
                        .map(|i| quote::format_ident!("field_{}", i))
                        .collect();

                    let field_inspections: Vec<_> = field_bindings.iter()
                        .enumerate()
                        .map(|(idx, binding)| {
                            quote! {
                                (
                                    inspect_core::FieldInfo::tuple(#idx),
                                    #binding.inspect(cx)
                                )
                            }
                        })
                        .collect();

                    quote! {
                        #name::#variant_name(#(#field_bindings),*) => {
                            cx.visit_node();
                            let variant = inspect_core::VariantInfo::new(#variant_name_str, #variant_idx);
                            let fields = vec![#(#field_inspections),*];

                            inspect_core::ValueRef::with_children(
                                inspect_core::Kind::Enum,
                                inspect_core::TypeInfo::new(#type_name),
                                inspect_core::Children::direct(fields),
                            )
                            .with_variant(variant)
                        }
                    }
                }
                Fields::Unit => {
                    quote! {
                        #name::#variant_name => {
                            cx.visit_node();
                            let variant = inspect_core::VariantInfo::new(#variant_name_str, #variant_idx);

                            inspect_core::ValueRef::with_children(
                                inspect_core::Kind::Enum,
                                inspect_core::TypeInfo::new(#type_name),
                                inspect_core::Children::direct(vec![]),
                            )
                            .with_variant(variant)
                        }
                    }
                }
            }
        })
        .collect();

    quote! {
        match self {
            #(#variant_arms),*
        }
    }
}

#[derive(Default)]
struct FieldAttributes {
    skip: bool,
    secret: bool,
    sensitive: bool,
    rename: Option<String>,
}

impl FieldAttributes {
    fn parse(attrs: &[syn::Attribute]) -> Self {
        let mut result = Self::default();

        for attr in attrs {
            if !attr.path().is_ident("inspect") {
                continue;
            }

            let _ = attr.parse_nested_meta(|meta| {
                if meta.path.is_ident("skip") {
                    result.skip = true;
                } else if meta.path.is_ident("secret") {
                    result.secret = true;
                } else if meta.path.is_ident("sensitive") {
                    result.sensitive = true;
                } else if meta.path.is_ident("rename") {
                    if let Ok(value) = meta.value() {
                        if let Ok(s) = value.parse::<syn::LitStr>() {
                            result.rename = Some(s.value());
                        }
                    }
                }
                Ok(())
            });
        }

        result
    }
}