extern crate proc_macro;
use proc_macro::TokenStream;
use proc_macro2::{Span, TokenStream as TokenStream2, TokenTree};
use proc_macro_error::{abort_call_site, proc_macro_error};
use quote::{format_ident, quote, ToTokens};
use std::{
fmt::{self, Display},
iter::FromIterator,
str::FromStr,
};
use syn::{
parse_macro_input, punctuated::Punctuated, Attribute, Data, DeriveInput, Expr, ExprLit, Fields,
Ident, Lit, Meta, MetaList, MetaNameValue, Path, PredicateType, Token, TraitBound,
TraitBoundModifiers, Type, TypeParamBound, WherePredicate,
};
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
enum FieldType {
Ordered,
Unordered,
UnorderedDelta,
Scalar,
Delta,
}
const VALID_FIELD_TYPES: &str =
"\"ordered\", \"unordered\", \"unordered-delta\", \"delta\", or \"scalar\"";
type Field = (String, Type, FieldType, String);
type ParsedField = (String, Type, ParsedAttrs);
type ParsedAttrs = Result<(Option<FieldType>, String), FieldTypeError>;
#[proc_macro_derive(Delta, attributes(delta_struct))]
#[proc_macro_error]
pub fn derive_delta(input: TokenStream) -> TokenStream {
let DeriveInput {
attrs,
vis,
ident,
mut generics,
data,
} = parse_macro_input!(input as DeriveInput);
let (default_field_type, delta_leader) =
match get_fieldtype_from_attrs(attrs.into_iter(), "default") {
Ok((v, delta_leader)) => (v.unwrap_or(FieldType::Scalar), delta_leader),
Err(e) => {
abort_call_site!(
"delta_struct(default = ...) for {} is not an accepted value, expected {}. {}",
ident,
VALID_FIELD_TYPES,
e,
);
}
};
let delta_leader = match proc_macro2::TokenStream::from_str(&delta_leader) {
Ok(v) => v,
Err(e) => {
abort_call_site!("error parsing delta leader as token stream {}", e);
}
};
let delta_ident = format_ident!("{}Delta", ident);
let og_where_clause = generics.where_clause.clone();
let ty_generics_only = generics.split_for_impl().1.to_token_stream();
let Generated {
delta_type,
output_ty,
delta_body,
apply_body,
} = match data {
Data::Struct(strukt) => struct_impl(
&ident,
&vis,
&delta_ident,
&delta_leader,
&generics,
&og_where_clause,
&ty_generics_only,
strukt.fields,
default_field_type,
),
Data::Enum(enom) => enum_impl(
&ident,
&vis,
&delta_ident,
&delta_leader,
&generics,
&og_where_clause,
&ty_generics_only,
enom.variants.into_iter().collect(),
default_field_type,
),
Data::Union(_) => {
abort_call_site!(
"delta_struct::Delta may only be derived for struct and enum types. {} is a union.",
ident
)
}
};
let partial_eq_types = generics
.type_params()
.map(|t| t.ident.clone())
.collect::<Vec<_>>();
let where_clause = generics.make_where_clause();
for ty in partial_eq_types {
let mut bounds = Punctuated::new();
let mut segments = Punctuated::new();
segments.push(Ident::new("std", Span::call_site()).into());
segments.push(Ident::new("cmp", Span::call_site()).into());
segments.push(Ident::new("PartialEq", Span::call_site()).into());
bounds.push(TypeParamBound::Trait(TraitBound {
paren_token: None,
modifiers: TraitBoundModifiers::default(),
lifetimes: None,
maybe: None,
path: Path {
leading_colon: Some(Token!(::)(Span::call_site())),
segments,
},
}));
where_clause
.predicates
.push(WherePredicate::Type(PredicateType {
attrs: Vec::new(),
lifetimes: None,
bounded_ty: Type::Verbatim(<Ident as Into<TokenTree>>::into(ty).into()),
colon_token: Token!(:)(Span::call_site()),
bounds,
}));
}
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
let delta_impl = quote! {
impl #impl_generics Delta for #ident #ty_generics #where_clause {
type Output = #output_ty;
fn delta(old: Self, new: Self) -> Option<Self::Output> {
#delta_body
}
#[allow(unreachable_patterns, unreachable_code)]
fn apply_delta(
&mut self,
delta: Self::Output,
) -> ::std::result::Result<(), ::delta_struct::Mismatch> {
#apply_body
}
}
};
let output = quote! {
#delta_type
#delta_impl
};
TokenStream::from(output)
}
struct Generated {
delta_type: TokenStream2,
output_ty: TokenStream2,
delta_body: TokenStream2,
apply_body: TokenStream2,
}
fn read_fields(owner: &Ident, fields: Fields, default_field_type: FieldType) -> (bool, Vec<Field>) {
let (named, collected) = match fields {
Fields::Named(named) => (
true,
collect_results(
named.named.into_iter().map(|field| {
(
field.ident.unwrap().to_string(),
field.ty,
get_fieldtype_from_attrs(field.attrs.into_iter(), "field_type"),
)
}),
default_field_type,
),
),
Fields::Unnamed(unnamed) => (
false,
collect_results(
unnamed.unnamed.into_iter().enumerate().map(|(i, field)| {
(
i.to_string(),
field.ty,
get_fieldtype_from_attrs(field.attrs.into_iter(), "field_type"),
)
}),
default_field_type,
),
),
Fields::Unit => (false, Ok(vec![])),
};
match collected {
Ok(fields) => (named, fields),
Err(bad_fields) => {
let bad_fields = format!("{:?}", bad_fields);
abort_call_site!(
"delta_struct(field_type = ...) for fields in {}: {} are not valid values. Expected {}.",
owner,
bad_fields,
VALID_FIELD_TYPES
)
}
}
}
#[allow(clippy::too_many_arguments)] fn struct_impl(
ident: &Ident,
vis: &syn::Visibility,
delta_ident: &Ident,
delta_leader: &TokenStream2,
generics: &syn::Generics,
og_where_clause: &Option<syn::WhereClause>,
ty_generics: &TokenStream2,
fields: Fields,
default_field_type: FieldType,
) -> Generated {
let (named, fields) = read_fields(ident, fields, default_field_type);
let delta_fields = delta_fields(named, fields.iter().cloned());
let (compute_let, compute_fields) =
delta_compute_fields(named, Source::Whole, fields.iter().cloned());
let (apply_let, apply_actions) = delta_apply_fields(named, Source::Whole, fields.into_iter());
let (delta_type, compute_init, apply_pattern) = if named {
(
quote! {
#delta_leader
#vis struct #delta_ident #generics #og_where_clause {
#delta_fields
}
},
quote!(Self::Output { #compute_fields }),
quote!(Self::Output { #apply_let }),
)
} else {
(
quote! {
#delta_leader
#vis struct #delta_ident #generics (#delta_fields) #og_where_clause;
},
quote!(#delta_ident(#compute_fields)),
quote!(#delta_ident(#apply_let)),
)
};
Generated {
delta_type,
output_ty: quote!(#delta_ident #ty_generics),
delta_body: quote! {
let mut delta_is_some = false;
#compute_let
if delta_is_some {
Some(#compute_init)
} else {
None
}
},
apply_body: quote! {
let #apply_pattern = delta;
#apply_actions
Ok(())
},
}
}
#[allow(clippy::too_many_arguments)] fn enum_impl(
ident: &Ident,
vis: &syn::Visibility,
delta_ident: &Ident,
delta_leader: &TokenStream2,
generics: &syn::Generics,
og_where_clause: &Option<syn::WhereClause>,
ty_generics: &TokenStream2,
variants: Vec<syn::Variant>,
default_field_type: FieldType,
) -> Generated {
if variants.is_empty() {
abort_call_site!(
"delta_struct::Delta cannot be derived for {}, which has no variants: an \
uninhabited type has no two values to differ.",
ident
)
}
let read: Vec<(Ident, bool, Vec<Field>)> = variants
.into_iter()
.map(|variant| {
let (named, fields) = read_fields(ident, variant.fields, default_field_type);
(variant.ident, named, fields)
})
.collect();
let mut delta_variants = Vec::new();
let mut diff_arms = Vec::new();
let mut apply_arms = Vec::new();
for (variant, named, fields) in &read {
if fields.is_empty() {
let pattern = variant_pattern("e!(Self), variant, *named, fields, Some("old"));
diff_arms.push(quote!((#pattern, Self::#variant) => None,));
continue;
}
let declared = delta_fields_inner(*named, false, fields.iter().cloned());
delta_variants.push(if *named {
quote!(#variant { #declared })
} else {
quote!(#variant(#declared))
});
let (compute_let, compute_fields) =
delta_compute_fields(*named, Source::Bound, fields.iter().cloned());
let old = variant_pattern("e!(Self), variant, *named, fields, Some("old"));
let new = variant_pattern("e!(Self), variant, *named, fields, Some("new"));
let init = if *named {
quote!(#delta_ident::#variant { #compute_fields })
} else {
quote!(#delta_ident::#variant(#compute_fields))
};
diff_arms.push(quote! {
(#old, #new) => {
let mut delta_is_some = false;
#compute_let
if delta_is_some {
Some(::delta_struct::EnumDelta::Delta(#init))
} else {
None
}
}
});
let (_, apply_actions) = delta_apply_fields(*named, Source::Bound, fields.iter().cloned());
let target = variant_pattern("e!(Self), variant, *named, fields, Some("self"));
let carried = variant_pattern("e!(#delta_ident), variant, *named, fields, None);
apply_arms.push(quote! {
(#target, #carried) => { #apply_actions }
});
}
let source_names = read
.iter()
.map(|(variant, ..)| quote!(Self::#variant { .. } => stringify!(#variant),));
let delta_names = read
.iter()
.filter(|(_, _, fields)| !fields.is_empty())
.map(|(variant, ..)| quote!(#delta_ident::#variant { .. } => stringify!(#variant),));
Generated {
delta_type: quote! {
#delta_leader
#vis enum #delta_ident #generics #og_where_clause {
#(#delta_variants,)*
}
},
output_ty: quote! {
::delta_struct::EnumDelta<#ident #ty_generics, #delta_ident #ty_generics>
},
delta_body: quote! {
#[allow(unreachable_patterns)] match (old, new) {
#(#diff_arms)*
(_, new) => Some(::delta_struct::EnumDelta::Became(new)),
}
},
apply_body: quote! {
let delta = match delta {
::delta_struct::EnumDelta::Became(new) => {
*self = new;
return Ok(());
}
::delta_struct::EnumDelta::Delta(delta) => delta,
};
match (&mut *self, delta) {
#(#apply_arms)*
(found, mismatched) => {
return Err(::delta_struct::Mismatch {
type_name: stringify!(#ident),
expected: match mismatched { #(#delta_names)* },
found: match found { #(#source_names)* },
});
}
}
Ok(())
},
}
}
fn variant_pattern(
path: &TokenStream2,
variant: &Ident,
named: bool,
fields: &[Field],
prefix: Option<&str>,
) -> TokenStream2 {
if fields.is_empty() {
return quote!(#path::#variant);
}
let bindings = fields
.iter()
.map(|(og_ident, ..)| {
let local = local_ident(named, og_ident);
match prefix {
Some(prefix) => format_ident!("{}_{}", prefix, local),
None => local,
}
})
.collect::<Vec<_>>();
if named {
let names = fields
.iter()
.map(|(og_ident, ..)| format_ident!("{}", og_ident));
quote!(#path::#variant { #(#names: #bindings),* })
} else {
quote!(#path::#variant( #(#bindings),* ))
}
}
fn delta_fields(named: bool, iter: impl Iterator<Item = Field>) -> proc_macro2::TokenStream {
delta_fields_inner(named, true, iter)
}
fn delta_fields_inner(
named: bool,
public: bool,
iter: impl Iterator<Item = Field>,
) -> proc_macro2::TokenStream {
let vis = public.then(|| quote!(pub));
FromIterator::from_iter(iter.map(|(ident, ty, field_ty, field_leader)| {
let field_leader = proc_macro2::TokenStream::from_str(&field_leader).unwrap();
let declared_ty = match field_ty {
FieldType::Ordered => {
quote!(::delta_struct::SeqDelta<<#ty as ::std::iter::IntoIterator>::Item>)
}
FieldType::Unordered => {
quote!(<#ty as ::delta_struct::Unordered>::Delta)
}
FieldType::UnorderedDelta => {
let entry = quote!(<#ty as ::std::iter::IntoIterator>::Item);
let key = quote!(<#entry as ::delta_struct::MapEntry>::Key);
let value = quote!(<#entry as ::delta_struct::MapEntry>::Value);
quote!(::delta_struct::MapDelta<#key, #value, <#value as Delta>::Output>)
}
FieldType::Scalar => quote!(::delta_struct::ScalarDelta<#ty>),
FieldType::Delta => quote!(::std::option::Option<<#ty as Delta>::Output>),
};
if named {
let ident = format_ident!("{}", ident);
quote! {
#field_leader
#vis #ident: #declared_ty,
}
} else {
quote! {
#field_leader
#vis #declared_ty,
}
}
}))
}
fn delta_compute_fields(
named: bool,
source: Source,
iter: impl Iterator<Item = Field>,
) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
iter.map(|(og_ident, ty, field_ty, _field_leader)| {
let ident = local_ident(named, &og_ident);
let (old, new) = source.sides(&og_ident, &ident);
let statements = match field_ty {
FieldType::Ordered | FieldType::UnorderedDelta => {
let module = collection_module(field_ty);
quote! {
let #ident = ::delta_struct::#module::diff(#old, #new);
delta_is_some = delta_is_some || !#ident.is_empty();
}
}
FieldType::Unordered => quote! {
let #ident = match <#ty as ::delta_struct::Unordered>::diff(#old, #new) {
Some(v) => {
delta_is_some = true;
v
}
None => ::std::default::Default::default(),
};
},
FieldType::Scalar => quote! {
let #ident = if #old != #new {
delta_is_some = true;
::delta_struct::ScalarDelta::Changed(#new)
} else {
::delta_struct::ScalarDelta::Unchanged
};
},
FieldType::Delta => quote! {
let #ident = Delta::delta(#old, #new);
delta_is_some = delta_is_some || #ident.is_some();
},
};
(statements, quote!(#ident,))
})
.unzip()
}
fn delta_apply_fields(
named: bool,
source: Source,
iter: impl Iterator<Item = Field>,
) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
iter.map(|(og_ident, ty, field_ty, _field_leader)| {
let ident = local_ident(named, &og_ident);
let target = source.target(&og_ident, &ident);
let statements = match field_ty {
FieldType::Ordered | FieldType::UnorderedDelta => {
let module = collection_module(field_ty);
let question = (field_ty == FieldType::UnorderedDelta).then(|| quote!(?));
quote! {
::delta_struct::#module::apply(&mut #target, #ident)#question;
}
}
FieldType::Unordered => quote! {
<#ty as ::delta_struct::Unordered>::apply(&mut #target, #ident);
},
FieldType::Scalar => quote! {
if let ::delta_struct::ScalarDelta::Changed(v) = #ident {
#target = v;
}
},
FieldType::Delta => quote! {
if let Some(v) = #ident {
#target.apply_delta(v)?;
}
},
};
(quote!(#ident,), statements)
})
.unzip()
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
enum Source {
Whole,
Bound,
}
impl Source {
fn sides(self, og_ident: &str, ident: &Ident) -> (TokenStream2, TokenStream2) {
match self {
Source::Whole => {
let og_ident = field_accessor(og_ident);
(quote!(old.#og_ident), quote!(new.#og_ident))
}
Source::Bound => {
let old = format_ident!("old_{}", ident);
let new = format_ident!("new_{}", ident);
(quote!(#old), quote!(#new))
}
}
}
fn target(self, og_ident: &str, ident: &Ident) -> TokenStream2 {
match self {
Source::Whole => {
let og_ident = field_accessor(og_ident);
quote!(self.#og_ident)
}
Source::Bound => {
let binding = format_ident!("self_{}", ident);
quote!((*#binding))
}
}
}
}
fn field_accessor(og_ident: &str) -> TokenStream2 {
FromStr::from_str(og_ident).unwrap()
}
fn local_ident(named: bool, og_ident: &str) -> Ident {
if named {
format_ident!("{}", og_ident)
} else {
format_ident!("field_{}", og_ident)
}
}
fn collection_module(field_ty: FieldType) -> Ident {
match field_ty {
FieldType::Ordered => format_ident!("seq"),
FieldType::UnorderedDelta => format_ident!("map"),
FieldType::Unordered | FieldType::Scalar | FieldType::Delta => {
unreachable!("{:?} does not have a fixed collection module", field_ty)
}
}
}
#[allow(clippy::manual_try_fold)] fn collect_results(
iter: impl Iterator<Item = ParsedField>,
default_field_type: FieldType,
) -> Result<Vec<Field>, Vec<String>> {
iter.fold(Ok(vec![]), |v, i| match (v, i) {
(Ok(mut v), (ident, b, Ok((c, d)))) => {
v.push((ident, b, c.unwrap_or(default_field_type), d));
Ok(v)
}
(Ok(_), (ident, _, Err(_))) => Err(vec![ident]),
(Err(mut v), (ident, _, Err(_))) => {
v.push(ident);
Err(v)
}
(v @ Err(_), _) => v,
})
}
enum FieldTypeError {
Syn(syn::Error),
UnrecognizedJunkFound(Vec<Meta>),
}
impl From<syn::Error> for FieldTypeError {
fn from(value: syn::Error) -> Self {
FieldTypeError::Syn(value)
}
}
impl Display for FieldTypeError {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match self {
FieldTypeError::Syn(error) => write!(f, "{error}"),
FieldTypeError::UnrecognizedJunkFound(metas) => {
let metas = metas
.iter()
.map(|m| m.into_token_stream().to_string())
.collect::<Vec<_>>()
.join(" ");
write!(
f,
"expected a comma separated list of named values, got {metas}"
)
}
}
}
}
#[allow(clippy::manual_try_fold)] fn get_fieldtype_from_attrs(iter: impl Iterator<Item = Attribute>, attr_name: &str) -> ParsedAttrs {
for attr in iter {
if let Meta::List(MetaList { path, .. }) = &attr.meta {
let Path { segments, .. } = path;
if segments
.iter()
.map(|p| &p.ident)
.eq(["delta_struct"].iter().cloned())
{
let nested =
attr.parse_args_with(Punctuated::<Meta, Token!(,)>::parse_terminated)?;
let values: Result<Vec<_>, Vec<Meta>> = nested
.iter()
.map(|meta| match meta {
Meta::NameValue(MetaNameValue {
path,
value:
Expr::Lit(ExprLit {
lit: Lit::Str(s), ..
}),
..
}) => Ok((path.get_ident().map(|i| i.to_string()), s.value())),
e => Err(e),
})
.fold(Ok(vec![]), |v, i| match (v, i) {
(Ok(mut v), Ok(i)) => {
v.push(i);
Ok(v)
}
(Ok(_), Err(e)) => Err(vec![e.clone()]),
(Err(mut v), Err(e)) => {
v.push(e.clone());
Err(v)
}
(v @ Err(_), _) => v,
});
let v = values.map_err(FieldTypeError::UnrecognizedJunkFound)?;
let mut field_type = None;
let mut delta_leader = String::new();
for i in v {
match i.0.as_deref() {
Some("delta_leader") => {
delta_leader = i.1;
}
a if Some(attr_name) == a => {
field_type = string_to_fieldtype(&i.1);
}
a => {
abort_call_site!("Unrecognized value {:?}", a);
}
}
}
return Ok((field_type, delta_leader));
}
}
}
Ok((None, String::new()))
}
fn string_to_fieldtype(s: &str) -> Option<FieldType> {
match s {
"ordered" => Some(FieldType::Ordered),
"unordered" => Some(FieldType::Unordered),
"unordered-delta" => Some(FieldType::UnorderedDelta),
"scalar" => Some(FieldType::Scalar),
"delta" => Some(FieldType::Delta),
_ => None,
}
}
#[proc_macro_derive(Fingerprint)]
#[proc_macro_error]
pub fn derive_fingerprint(input: TokenStream) -> TokenStream {
let DeriveInput {
ident,
mut generics,
data,
..
} = parse_macro_input!(input as DeriveInput);
let body = match data {
Data::Struct(strukt) => {
fingerprint_calls(strukt.fields.iter().enumerate().map(|(i, field)| {
match &field.ident {
Some(ident) => quote!(self.#ident),
None => {
let index = syn::Index::from(i);
quote!(self.#index)
}
}
}))
}
Data::Enum(enom) => {
let arms = enom.variants.into_iter().enumerate().map(|(index, variant)| {
let variant_ident = variant.ident;
let bindings = binding_idents(&variant.fields);
let pattern = match &variant.fields {
Fields::Named(_) => quote!(Self::#variant_ident { #(#bindings),* }),
Fields::Unnamed(_) => quote!(Self::#variant_ident( #(#bindings),* )),
Fields::Unit => quote!(Self::#variant_ident),
};
let fields = fingerprint_calls(bindings.iter().map(|b| quote!(#b)));
let index = index as u32;
quote! {
#pattern => {
::delta_struct::Fingerprint::fingerprint(&#index, hasher);
#fields
}
}
});
quote! {
match self {
#(#arms)*
}
}
}
_ => abort_call_site!(
"delta_struct::Fingerprint may only be derived for struct and enum types. {} is neither.",
ident
),
};
let fingerprint_types = generics
.type_params()
.map(|t| t.ident.clone())
.collect::<Vec<_>>();
let where_clause = generics.make_where_clause();
for ty in fingerprint_types {
where_clause
.predicates
.push(syn::parse_quote!(#ty: ::delta_struct::Fingerprint));
}
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
TokenStream::from(quote! {
impl #impl_generics ::delta_struct::Fingerprint for #ident #ty_generics #where_clause {
fn fingerprint(&self, hasher: &mut ::delta_struct::fingerprint::Hasher) {
#body
}
}
})
}
fn binding_idents(fields: &Fields) -> Vec<Ident> {
fields
.iter()
.enumerate()
.map(|(i, field)| match &field.ident {
Some(ident) => ident.clone(),
None => format_ident!("field_{}", i),
})
.collect()
}
fn fingerprint_calls(
exprs: impl Iterator<Item = proc_macro2::TokenStream>,
) -> proc_macro2::TokenStream {
let calls = exprs.map(|expr| quote!(::delta_struct::Fingerprint::fingerprint(&#expr, hasher);));
quote!(#(#calls)*)
}