use proc_macro::TokenStream;
use quote::{format_ident, quote, quote_spanned};
use syn::{
AngleBracketedGenericArguments, Data, DeriveInput, Fields, GenericArgument, PathArguments,
PathSegment, Type, parse_macro_input,
};
use super::derive_fields::parse_nested_attr;
fn is_std_string(ty: &Type) -> bool {
matches!(
ty,
Type::Path(tp)
if tp.qself.is_none()
&& tp.path.segments.last().is_some_and(|seg| seg.ident == "String" && tp.path.segments.len() == 1)
)
}
enum Container<'a> {
List(&'a Type),
Set(&'a Type),
Map(&'a Type, &'a Type),
}
fn container_kind(ty: &Type) -> Option<Container<'_>> {
let Type::Path(tp) = ty else { return None };
let seg = tp.path.segments.last()?;
match seg.ident.to_string().as_str() {
"Vec" => {
if let PathArguments::AngleBracketed(AngleBracketedGenericArguments { args, .. }) =
&seg.arguments
{
if let Some(GenericArgument::Type(inner)) = args.first() {
return Some(Container::List(inner));
}
}
}
"HashSet" => {
if let PathArguments::AngleBracketed(AngleBracketedGenericArguments { args, .. }) =
&seg.arguments
{
if let Some(GenericArgument::Type(inner)) = args.first() {
return Some(Container::Set(inner));
}
}
}
"HashMap" => {
if let PathArguments::AngleBracketed(AngleBracketedGenericArguments { args, .. }) =
&seg.arguments
{
let mut it = args.iter();
if let (Some(GenericArgument::Type(k)), Some(GenericArgument::Type(v))) =
(it.next(), it.next())
{
return Some(Container::Map(k, v));
}
}
}
_ => {}
}
None
}
pub fn derive_changes_impl(input: TokenStream) -> TokenStream {
let DeriveInput {
ident,
data,
generics,
..
} = parse_macro_input!(input as DeriveInput);
let Data::Struct(data_struct) = data else {
return syn::Error::new_spanned(ident, "Diff can only be derived for structs")
.to_compile_error()
.into();
};
let Fields::Named(fields_named) = data_struct.fields else {
return syn::Error::new_spanned(ident, "Diff needs named fields")
.to_compile_error()
.into();
};
let enum_ident = format_ident!("{ident}Change");
let snapshot_ident = format_ident!("{ident}Snapshot");
let (_, ty_generics, where_clause) = generics.split_for_impl();
let enum_lifetime = quote!('a);
let mut snapshot_fields = Vec::new();
let mut snapshot_inits = Vec::new();
for field in &fields_named.named {
let field_ident = field.ident.as_ref().unwrap();
let ty = &field.ty;
let borrowed_ty = if is_std_string(ty) {
quote_spanned! { field_ident.span() => ::std::borrow::Cow<'a, str> }
} else {
quote_spanned! { field_ident.span() => &'a #ty }
};
snapshot_fields.push(quote_spanned! { field_ident.span() =>
#field_ident : #borrowed_ty
});
let init_expr = if is_std_string(ty) {
quote_spanned! { field_ident.span() => ::std::borrow::Cow::Borrowed(src.#field_ident.as_str()) }
} else {
quote_spanned! { field_ident.span() => &src.#field_ident }
};
snapshot_inits.push(quote_spanned! { field_ident.span() =>
#field_ident : #init_expr
});
}
let snapshot_struct = quote! {
#[allow(non_camel_case_types)]
#[derive(Debug, Clone)]
pub struct #snapshot_ident<'a> { #( #snapshot_fields, )* }
impl<'a> From<&'a #ident #ty_generics> for #snapshot_ident<'a> {
fn from(src: &'a #ident #ty_generics) -> Self {
Self { #( #snapshot_inits, )* }
}
}
};
let mut enum_variants = Vec::new();
let mut diff_arms = Vec::new();
let self_variant_ident = format_ident!("self_");
enum_variants.push(
quote_spanned! {ident.span() => #self_variant_ident(#snapshot_ident<#enum_lifetime>) },
);
diff_arms.push(quote_spanned! {ident.span() =>
if old != new {
out.push(#enum_ident::#self_variant_ident(#snapshot_ident::from(new)));
}
});
for field in &fields_named.named {
let field_ident = field.ident.as_ref().unwrap();
let ty = &field.ty;
let span = field_ident.span();
let nested_attr = field.attrs.iter().find_map(parse_nested_attr);
if nested_attr.is_some() {
if let Type::Path(tp) = ty {
let nested_change_ident = tp
.path
.segments
.last()
.map(|PathSegment { ident, .. }| format_ident!("{ident}Change"))
.unwrap();
enum_variants.push(quote_spanned! { span =>
#field_ident(#nested_change_ident<#enum_lifetime>)
});
diff_arms.push(quote_spanned! { span =>
{
let mut __subs = Vec::new();
<#ty as ::differ::HasChanges>::collect_changes(
&old.#field_ident,
&new.#field_ident,
&mut __subs,
);
for __sub in __subs {
out.push(#enum_ident::#field_ident(__sub));
}
}
});
}
continue;
}
if let Some(kind) = container_kind(ty) {
match kind {
Container::List(elem_ty) => {
let changed_ty =
quote_spanned! { span => ::differ::Changed<#enum_lifetime, #elem_ty> };
enum_variants.push(quote_spanned! { span => #field_ident(#changed_ty) });
diff_arms.push(quote_spanned! { span =>
{
use ::std::collections::HashSet;
let __old: HashSet<_> = old.#field_ident.iter().collect();
let __new: HashSet<_> = new.#field_ident.iter().collect();
for __item in __old.difference(&__new) {
out.push(#enum_ident::#field_ident(::differ::Changed::Removed(*__item)));
}
for __item in __new.difference(&__old) {
out.push(#enum_ident::#field_ident(::differ::Changed::Added(*__item)));
}
}
});
}
Container::Set(elem_ty) => {
let changed_ty =
quote_spanned! { span => ::differ::Changed<#enum_lifetime, #elem_ty> };
enum_variants.push(quote_spanned! { span => #field_ident(#changed_ty) });
diff_arms.push(quote_spanned! { span =>
{
for __item in old.#field_ident.difference(&new.#field_ident) {
out.push(#enum_ident::#field_ident(::differ::Changed::Removed(__item)));
}
for __item in new.#field_ident.difference(&old.#field_ident) {
out.push(#enum_ident::#field_ident(::differ::Changed::Added(__item)));
}
}
});
}
Container::Map(key_ty, val_ty) => {
let changed_ty =
quote! { ::differ::Changed<#enum_lifetime, (&#key_ty, &#val_ty)> };
enum_variants.push(quote_spanned! { span => #field_ident(#changed_ty) });
diff_arms.push(quote_spanned! { span =>
{
for (__k, __v) in &old.#field_ident {
if !new.#field_ident.contains_key(__k) {
out.push(#enum_ident::#field_ident(::differ::Changed::Removed((__k, __v))));
}
}
for (__k, __v) in &new.#field_ident {
if !old.#field_ident.contains_key(__k) {
out.push(#enum_ident::#field_ident(::differ::Changed::Added((__k, __v))));
}
}
}
});
}
}
continue;
}
let scalar_ty = if is_std_string(ty) {
quote_spanned! { span => ::std::borrow::Cow<#enum_lifetime, str> }
} else {
quote_spanned! { span => &#enum_lifetime #ty }
};
enum_variants.push(quote_spanned! { span => #field_ident(#scalar_ty) });
let val_expr = if is_std_string(ty) {
quote_spanned! { span => ::std::borrow::Cow::Borrowed(new.#field_ident.as_str()) }
} else {
quote_spanned! { span => &new.#field_ident }
};
diff_arms.push(quote_spanned! { span =>
if old.#field_ident != new.#field_ident {
out.push(#enum_ident::#field_ident(#val_expr));
}
});
}
let expanded = quote_spanned! { ident.span() =>
#snapshot_struct
#[allow(non_camel_case_types)]
#[derive(Debug)]
pub enum #enum_ident<#enum_lifetime> { #( #enum_variants, )* }
impl ::differ::HasChanges for #ident #ty_generics #where_clause {
type Change<'a> = #enum_ident<'a> where Self: 'a;
fn collect_changes<'a>(
old: &'a Self,
new: &'a Self,
out: &mut Vec<Self::Change<'a>>,
) where Self: 'a {
#(#diff_arms)*
}
}
};
TokenStream::from(expanded)
}