Skip to main content

aura_anim_macros/
lib.rs

1//! Derive macros for Aura animation values.
2
3use proc_macro::TokenStream;
4use proc_macro_crate::{FoundCrate, crate_name};
5use quote::{format_ident, quote};
6use syn::{
7    Attribute, Data, DeriveInput, Fields, Generics, Ident, Member, Path, Type, Visibility,
8    parse_macro_input, parse_quote,
9};
10
11/// Derives field-by-field interpolation for a struct.
12///
13/// Every field must implement `Animatable`. Named, tuple, and unit structs are
14/// supported.
15#[proc_macro_derive(Animatable, attributes(animatable))]
16pub fn derive_animatable(input: TokenStream) -> TokenStream {
17    let input = parse_macro_input!(input as DeriveInput);
18    expand(input)
19        .unwrap_or_else(syn::Error::into_compile_error)
20        .into()
21}
22
23/// Creates a type-safe descriptor for a struct field.
24///
25/// The input uses field access syntax, for example `field!(Position::x)` or
26/// `field!(Offset::0)`.
27#[proc_macro]
28pub fn field(input: TokenStream) -> TokenStream {
29    expand_field(input.into())
30        .unwrap_or_else(syn::Error::into_compile_error)
31        .into()
32}
33
34fn expand(input: DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
35    let path = crate_path();
36    let name = input.ident;
37    let visibility = input.vis;
38    let generated_fields_type = fields_name(&input.attrs, &name)?;
39    let Data::Struct(data) = input.data else {
40        return Err(syn::Error::new_spanned(
41            name,
42            "Animatable can only be derived for structs",
43        ));
44    };
45
46    let descriptor_generics = input.generics.clone();
47    let mut interpolation_generics = input.generics;
48    let field_types = data
49        .fields
50        .iter()
51        .map(|field| field.ty.clone())
52        .collect::<Vec<_>>();
53    let where_clause = interpolation_generics.make_where_clause();
54    for field_type in &field_types {
55        where_clause
56            .predicates
57            .push(parse_quote!(#field_type: #path::Animatable));
58    }
59    let (impl_generics, type_generics, where_clause) = interpolation_generics.split_for_impl();
60
61    let interpolate_body = match &data.fields {
62        Fields::Named(fields) => {
63            let names = fields
64                .named
65                .iter()
66                .map(|field| field.ident.as_ref().unwrap())
67                .collect::<Vec<_>>();
68            quote! {
69                Self {
70                    #(
71                        #names: #path::Interpolate::interpolate_progress(
72                            &from.#names,
73                            &to.#names,
74                            progress,
75                        )
76                    ),*
77                }
78            }
79        }
80        Fields::Unnamed(fields) => {
81            let indexes = (0..fields.unnamed.len())
82                .map(syn::Index::from)
83                .collect::<Vec<_>>();
84            quote! {
85                Self(
86                    #(
87                        #path::Interpolate::interpolate_progress(
88                            &from.#indexes,
89                            &to.#indexes,
90                            progress,
91                        )
92                    ),*
93                )
94            }
95        }
96        Fields::Unit => quote!(Self),
97    };
98
99    let field_descriptors = expand_field_descriptors(
100        &path,
101        &name,
102        &visibility,
103        &descriptor_generics,
104        &data.fields,
105        &generated_fields_type,
106    );
107
108    let interpolate_impl = quote! {
109        impl #impl_generics #path::Interpolate for #name #type_generics #where_clause {
110            fn interpolate_progress(
111                from: &Self,
112                to: &Self,
113                progress: #path::InterpolationProgress,
114            ) -> Self {
115                #interpolate_body
116            }
117        }
118    };
119
120    Ok(quote! {
121        #interpolate_impl
122        #field_descriptors
123    })
124}
125
126fn expand_field_descriptors(
127    path: &proc_macro2::TokenStream,
128    struct_name: &Ident,
129    visibility: &Visibility,
130    generics: &Generics,
131    fields: &Fields,
132    generated_type: &Ident,
133) -> proc_macro2::TokenStream {
134    let (impl_generics, type_generics, where_clause) = generics.split_for_impl();
135    let struct_type = quote!(#struct_name #type_generics);
136    let descriptor_constants = match fields {
137        Fields::Named(fields) => fields
138            .named
139            .iter()
140            .map(|field| {
141                let field_visibility = &field.vis;
142                let member = field.ident.as_ref().expect("named field has an identifier");
143                let field_type = &field.ty;
144
145                quote! {
146                    #[allow(non_upper_case_globals)]
147                    #field_visibility const #member: #path::Field<#struct_type, #field_type> =
148                        #path::Field::new(
149                            stringify!(#member),
150                            |value: &#struct_type| &value.#member,
151                            |value: &mut #struct_type| &mut value.#member,
152                        );
153                }
154            })
155            .collect::<Vec<_>>(),
156        Fields::Unnamed(fields) => fields
157            .unnamed
158            .iter()
159            .enumerate()
160            .map(|(index, field)| {
161                let field_visibility = &field.vis;
162                let descriptor_name = format_ident!("_{index}");
163                let field_index = syn::Index::from(index);
164                let field_type = &field.ty;
165
166                quote! {
167                    #[allow(non_upper_case_globals)]
168                    #field_visibility const #descriptor_name: #path::Field<#struct_type, #field_type> =
169                        #path::Field::new(
170                            stringify!(#field_index),
171                            |value: &#struct_type| &value.#field_index,
172                            |value: &mut #struct_type| &mut value.#field_index,
173                        );
174                }
175            })
176            .collect::<Vec<_>>(),
177        Fields::Unit => Vec::new(),
178    };
179
180    quote! {
181        #[doc = concat!("Field descriptors generated for [`", stringify!(#struct_name), "`].")]
182        #visibility struct #generated_type #generics {
183            __aura_anim_marker: ::core::marker::PhantomData<fn() -> #struct_type>,
184        }
185
186        impl #impl_generics #generated_type #type_generics #where_clause
187        {
188            #(#descriptor_constants)*
189        }
190    }
191}
192
193fn fields_name(attributes: &[Attribute], struct_name: &Ident) -> syn::Result<Ident> {
194    let mut name = format_ident!("{struct_name}Fields");
195    let mut configured = false;
196
197    for attribute in attributes {
198        if !attribute.path().is_ident("animatable") {
199            continue;
200        }
201
202        attribute.parse_nested_meta(|meta| {
203            if !meta.path.is_ident("fields") {
204                return Err(meta.error("unsupported animatable option"));
205            }
206            if configured {
207                return Err(meta.error("field descriptor name was already configured"));
208            }
209
210            let path: Path = meta.value()?.parse()?;
211            let Some(identifier) = path.get_ident() else {
212                return Err(meta.error("field descriptor name must be a single identifier"));
213            };
214            name = identifier.clone();
215            configured = true;
216            Ok(())
217        })?;
218    }
219
220    Ok(name)
221}
222
223fn expand_field(input: proc_macro2::TokenStream) -> syn::Result<proc_macro2::TokenStream> {
224    let tokens = input.into_iter().collect::<Vec<_>>();
225    let separator = tokens
226        .windows(2)
227        .enumerate()
228        .filter(|(_, pair)| {
229            matches!(&pair[0], proc_macro2::TokenTree::Punct(punct) if punct.as_char() == ':')
230                && matches!(&pair[1], proc_macro2::TokenTree::Punct(punct) if punct.as_char() == ':')
231        })
232        .map(|(index, _)| index)
233        .next_back()
234        .ok_or_else(|| {
235            syn::Error::new(
236                proc_macro2::Span::call_site(),
237                "expected a field path such as Position::x",
238            )
239        })?;
240
241    let field_type = syn::parse2::<Type>(tokens[..separator].iter().cloned().collect())?;
242    let member = syn::parse2::<Member>(tokens[separator + 2..].iter().cloned().collect())?;
243    let path = crate_path();
244    let name = match &member {
245        Member::Named(identifier) => identifier.to_string(),
246        Member::Unnamed(index) => index.index.to_string(),
247    };
248
249    Ok(quote! {
250        #path::Field::new(
251            #name,
252            |value: &#field_type| &value.#member,
253            |value: &mut #field_type| &mut value.#member,
254        )
255    })
256}
257
258fn crate_path() -> proc_macro2::TokenStream {
259    for package in ["aura-anim-core", "aura-anim"] {
260        if let Ok(found) = crate_name(package) {
261            return match found {
262                FoundCrate::Itself if package == "aura-anim-core" => quote!(crate),
263                FoundCrate::Itself => quote!(::aura_anim),
264                FoundCrate::Name(name) => {
265                    let ident = format_ident!("{name}");
266                    quote!(::#ident)
267                }
268            };
269        }
270    }
271
272    quote!(::aura_anim)
273}