use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::{DeriveInput, Fields, Index};
pub fn expand(item: DeriveInput) -> Result<TokenStream, syn::Error> {
let input_name = &item.ident;
let data = match &item.data {
syn::Data::Struct(s) => s,
_ => {
return Err(syn::Error::new_spanned(
&item,
"drv::Input can only be derived on structs",
));
}
};
let snapshot_ident = format_ident!("__Drv{}", input_name);
let fields: Vec<FieldInfo> = match &data.fields {
Fields::Named(f) => f
.named
.iter()
.map(|field| FieldInfo {
accessor: {
let id = field.ident.as_ref().unwrap();
quote! { #id }
},
decl_head: {
let attrs = &field.attrs;
let vis = &field.vis;
let id = field.ident.as_ref().unwrap();
quote! { #(#attrs)* #vis #id: }
},
build_head: {
let id = field.ident.as_ref().unwrap();
quote! { #id: }
},
ty: field.ty.clone(),
})
.collect(),
Fields::Unnamed(f) => f
.unnamed
.iter()
.enumerate()
.map(|(i, field)| {
let idx = Index::from(i);
FieldInfo {
accessor: quote! { #idx },
decl_head: {
let attrs = &field.attrs;
let vis = &field.vis;
quote! { #(#attrs)* #vis }
},
build_head: quote! {},
ty: field.ty.clone(),
}
})
.collect(),
Fields::Unit => Vec::new(),
};
let mut snap_decls = Vec::new();
let mut eq_checks = Vec::new();
let mut snap_stores = Vec::new();
for f in &fields {
if is_phantom_data(&f.ty) {
continue;
}
let fty_static = type_to_static_form(&f.ty);
let decl_head = &f.decl_head;
let build_head = &f.build_head;
let accessor = &f.accessor;
snap_decls.push(quote! {
#decl_head <#fty_static as ::drv::ToStatic>::Static
});
snap_stores.push(quote! {
#build_head ::drv::ToStatic::to_static(&self.#accessor)
});
eq_checks.push(quote! {
::drv::ToStatic::eq_static(&self.#accessor, &other.#accessor)
});
}
let eq_body = if eq_checks.is_empty() {
quote! { true }
} else {
quote! { #(#eq_checks)&&* }
};
let (shadow_generics, shadow_ty_args) = shadow_struct_generics(&item.generics);
let shadow_where = build_shadow_where_clause(&item.generics);
let (snap_struct_with_generics, snap_build_expr) = match &data.fields {
Fields::Named(_) => (
quote! {
pub struct #snapshot_ident #shadow_generics #shadow_where {
#(#snap_decls,)*
}
},
quote! {
#snapshot_ident {
#(#snap_stores,)*
}
},
),
Fields::Unnamed(_) => (
quote! {
pub struct #snapshot_ident #shadow_generics (
#(#snap_decls,)*
) #shadow_where;
},
quote! {
#snapshot_ident(
#(#snap_stores,)*
)
},
),
Fields::Unit => (
quote! {
pub struct #snapshot_ident #shadow_generics #shadow_where;
},
quote! { #snapshot_ident },
),
};
let (impl_generics, ty_generics, where_clause) = item.generics.split_for_impl();
let extra_static_bounds = extra_static_bounds_for_type_params(&item.generics);
let mut all_preds: Vec<TokenStream> = Vec::new();
if let Some(w) = where_clause {
for p in &w.predicates {
all_preds.push(quote! { #p });
}
}
all_preds.extend(extra_static_bounds);
let impl_where = if all_preds.is_empty() {
quote! {}
} else {
quote! { where #(#all_preds),* }
};
Ok(quote! {
#[doc(hidden)]
#[allow(non_camel_case_types)]
#snap_struct_with_generics
impl #impl_generics ::drv::ToStatic for #input_name #ty_generics #impl_where {
type Static = #snapshot_ident #shadow_ty_args;
fn to_static(&self) -> Self::Static {
#snap_build_expr
}
fn eq_static(&self, other: &Self::Static) -> bool {
#eq_body
}
}
})
}
fn shadow_struct_generics(generics: &syn::Generics) -> (TokenStream, TokenStream) {
let mut decl_params: Vec<TokenStream> = Vec::new();
let mut use_args: Vec<TokenStream> = Vec::new();
for p in &generics.params {
match p {
syn::GenericParam::Lifetime(_) => {} syn::GenericParam::Type(t) => {
let id = &t.ident;
let bounds = &t.bounds;
let decl = if bounds.is_empty() {
quote! { #id: 'static }
} else {
quote! { #id: 'static + #bounds }
};
decl_params.push(decl);
use_args.push(quote! { #id });
}
syn::GenericParam::Const(c) => {
let id = &c.ident;
let ty = &c.ty;
decl_params.push(quote! { const #id: #ty });
use_args.push(quote! { #id });
}
}
}
let decl = if decl_params.is_empty() {
quote! {}
} else {
quote! { <#(#decl_params),*> }
};
let use_ = if use_args.is_empty() {
quote! {}
} else {
quote! { <#(#use_args),*> }
};
(decl, use_)
}
fn build_shadow_where_clause(generics: &syn::Generics) -> TokenStream {
let Some(w) = &generics.where_clause else {
return quote! {};
};
let dropped_lifetimes: Vec<syn::Ident> = generics
.params
.iter()
.filter_map(|p| match p {
syn::GenericParam::Lifetime(lt) => Some(lt.lifetime.ident.clone()),
_ => None,
})
.collect();
let preds: Vec<TokenStream> = w
.predicates
.iter()
.filter(|p| match p {
syn::WherePredicate::Lifetime(l) => !dropped_lifetimes.contains(&l.lifetime.ident),
_ => true,
})
.map(|p| quote! { #p })
.collect();
if preds.is_empty() {
quote! {}
} else {
quote! { where #(#preds),* }
}
}
fn extra_static_bounds_for_type_params(generics: &syn::Generics) -> Vec<TokenStream> {
generics
.params
.iter()
.filter_map(|p| match p {
syn::GenericParam::Type(t) => {
let id = &t.ident;
Some(quote! { #id: 'static })
}
_ => None,
})
.collect()
}
struct FieldInfo {
accessor: TokenStream,
decl_head: TokenStream,
build_head: TokenStream,
ty: syn::Type,
}
fn is_phantom_data(ty: &syn::Type) -> bool {
if let syn::Type::Path(p) = ty {
if let Some(last) = p.path.segments.last() {
return last.ident == "PhantomData";
}
}
false
}
fn type_to_static_form(ty: &syn::Type) -> syn::Type {
let stripped = match ty {
syn::Type::Reference(r) => (*r.elem).clone(),
other => other.clone(),
};
let mut out = stripped;
substitute_lifetimes_static(&mut out);
out
}
fn substitute_lifetimes_static(ty: &mut syn::Type) {
use syn::{GenericArgument, PathArguments, Type};
let static_lt = syn::Lifetime::new("'static", proc_macro2::Span::call_site());
match ty {
Type::Reference(r) => {
r.lifetime = Some(static_lt.clone());
substitute_lifetimes_static(&mut r.elem);
}
Type::Path(p) => {
for seg in &mut p.path.segments {
if let PathArguments::AngleBracketed(args) = &mut seg.arguments {
for arg in &mut args.args {
match arg {
GenericArgument::Lifetime(lt) => {
*lt = static_lt.clone();
}
GenericArgument::Type(t) => substitute_lifetimes_static(t),
_ => {}
}
}
}
}
}
Type::Tuple(t) => {
for elem in &mut t.elems {
substitute_lifetimes_static(elem);
}
}
Type::Array(a) => substitute_lifetimes_static(&mut a.elem),
Type::Slice(s) => substitute_lifetimes_static(&mut s.elem),
Type::Ptr(p) => substitute_lifetimes_static(&mut p.elem),
Type::Paren(p) => substitute_lifetimes_static(&mut p.elem),
Type::Group(g) => substitute_lifetimes_static(&mut g.elem),
_ => {}
}
}