Skip to main content

lattices_macro/
lib.rs

1//! Macros for the `lattices` crate.
2//!
3//! See [`[derive(Lattice)`](Lattice).
4#![warn(missing_docs)]
5
6use proc_macro2::{Span, TokenStream};
7use quote::{format_ident, quote};
8use syn::punctuated::Punctuated;
9use syn::visit_mut::VisitMut;
10use syn::{
11    Field, FieldsNamed, FieldsUnnamed, Generics, Ident, Index, ItemStruct, Member, Token,
12    WhereClause, WherePredicate, parse_macro_input,
13};
14
15/// Tokens to reference the `lattices` crate.
16fn root() -> TokenStream {
17    use std::env::{VarError, var as env_var};
18
19    use proc_macro_crate::FoundCrate;
20
21    if matches!(
22        proc_macro_crate::crate_name("lattices_macro"),
23        Ok(FoundCrate::Itself)
24    ) {
25        return quote! { lattices };
26    }
27
28    let lattices_crate_name = env!("CARGO_PKG_NAME").strip_suffix("_macro").unwrap();
29    let lattices_crate_ident = lattices_crate_name.replace('-', "_");
30    let lattices_crate = proc_macro_crate::crate_name(lattices_crate_name)
31        .unwrap_or_else(|_| panic!("`{lattices_crate_name}` should be present in `Cargo.toml`"));
32    match lattices_crate {
33        FoundCrate::Itself => {
34            if Err(VarError::NotPresent) == env_var("CARGO_BIN_NAME")
35                && Ok(&*lattices_crate_ident) == env_var("CARGO_CRATE_NAME").as_deref()
36            {
37                // In the crate itself, including unit tests.
38                quote! { crate }
39            } else {
40                // In an integration test, example, bench, etc.
41                let ident = Ident::new(&lattices_crate_ident, Span::call_site());
42                quote! { ::#ident }
43            }
44        }
45        FoundCrate::Name(name) => {
46            let ident = Ident::new(&name, Span::call_site());
47            quote! { ::#ident }
48        }
49    }
50}
51
52/// Renames the generics and returns the updated `WherePredicate`s.
53fn rename_generics(
54    item_struct: &mut ItemStruct,
55    rename: impl FnMut(&Ident) -> Ident,
56) -> Vec<WherePredicate> {
57    struct RenameGenerics<F> {
58        rename: F,
59        names: Vec<Ident>,
60        pub triggered: bool,
61    }
62    impl<F> VisitMut for RenameGenerics<F>
63    where
64        F: FnMut(&Ident) -> Ident,
65    {
66        fn visit_ident_mut(&mut self, i: &mut Ident) {
67            if self.names.contains(i) {
68                *i = (self.rename)(i);
69                self.triggered = true;
70            }
71        }
72    }
73
74    let names = item_struct
75        .generics
76        .type_params()
77        .map(|type_param| type_param.ident.clone())
78        .collect();
79    let mut visit = RenameGenerics {
80        rename,
81        names,
82        triggered: false,
83    };
84
85    let mut out = Vec::new();
86    if let Some(where_clause) = &mut item_struct.generics.where_clause {
87        for where_predicate in where_clause.predicates.iter_mut() {
88            visit.visit_where_predicate_mut(where_predicate);
89            if std::mem::take(&mut visit.triggered) {
90                out.push(where_predicate.clone());
91            }
92        }
93    }
94    for type_param in item_struct.generics.type_params_mut() {
95        visit.visit_type_param_mut(type_param);
96    }
97    for field in item_struct.fields.iter_mut() {
98        visit.visit_type_mut(&mut field.ty);
99    }
100    out
101}
102
103/// Ensures that `punctuated` has trailing punctuation (or is empty).
104fn ensure_trailing<T, P>(punctuated: &mut Punctuated<T, P>)
105where
106    P: Default,
107{
108    if !punctuated.empty_or_trailing() {
109        punctuated.push_punct(Default::default());
110    }
111}
112
113#[doc = include_str!("../README.md")]
114#[proc_macro_derive(Lattice)]
115pub fn derive_lattice_macro(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
116    derive_lattice(&process_item_struct(parse_macro_input!(item))).into()
117}
118/// Derives lattice `Merge`.
119///
120/// See [`#[derive(Lattice)]`](`Lattice`) for more info.
121#[proc_macro_derive(Merge)]
122pub fn derive_merge_macro(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
123    derive_merge(&process_item_struct(parse_macro_input!(item))).into()
124}
125/// Derives [`PartialEq`], [`PartialOrd`], and `LatticeOrd` together.
126///
127/// See [`#[derive(Lattice)]`](`Lattice`) for more info.
128#[proc_macro_derive(LatticeOrd)]
129pub fn derive_lattice_ord_macro(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
130    derive_lattice_ord(&process_item_struct(parse_macro_input!(item))).into()
131}
132/// Derives lattice `IsBot`.
133///
134/// See [`#[derive(Lattice)]`](`Lattice`) for more info.
135#[proc_macro_derive(IsBot)]
136pub fn derive_is_bot_macro(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
137    derive_is_bot(&process_item_struct(parse_macro_input!(item))).into()
138}
139/// Derives lattice `IsTop`.
140///
141/// See [`#[derive(Lattice)]`](`Lattice`) for more info.
142#[proc_macro_derive(IsTop)]
143pub fn derive_is_top_macro(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
144    derive_is_top(&process_item_struct(parse_macro_input!(item))).into()
145}
146/// Derives `LatticeFrom`.
147///
148/// See [`#[derive(Lattice)]`](`Lattice`) for more info.
149#[proc_macro_derive(LatticeFrom)]
150pub fn derive_lattice_from_macro(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
151    derive_lattice_from(&process_item_struct(parse_macro_input!(item))).into()
152}
153
154/// [`process_item_struct`] return value helper struct.
155struct ProcessItemStruct {
156    root: TokenStream,
157    item_struct: ItemStruct,
158    item_struct_renamed: ItemStruct,
159    self_where_predicates: Punctuated<WherePredicate, Token![,]>,
160    both_where_predicates: Punctuated<WherePredicate, Token![,]>,
161    field_names: Vec<Member>,
162    combined_generics: Generics,
163}
164/// Helper for common pre-processing code shared between macros.
165fn process_item_struct(item_struct: ItemStruct) -> ProcessItemStruct {
166    let mut item_struct_renamed = item_struct.clone();
167    let extra_where_predicates = rename_generics(&mut item_struct_renamed, |ident| {
168        format_ident!("__{}Other", ident)
169    });
170
171    // Basic `where` predicates, no extras.
172    let mut self_where_predicates = item_struct
173        .generics
174        .where_clause
175        .clone()
176        .map(|WhereClause { predicates, .. }| predicates)
177        .unwrap_or_default();
178    ensure_trailing(&mut self_where_predicates);
179    // Basic `where` predicates for combined original and renamed parameters.
180    let mut both_where_predicates = self_where_predicates.clone();
181    both_where_predicates.extend(extra_where_predicates);
182    ensure_trailing(&mut both_where_predicates);
183
184    // Fields.
185    let field_names = match &item_struct.fields {
186        syn::Fields::Named(FieldsNamed { named, .. }) => named
187            .iter()
188            .map(|Field { ident, .. }| Member::Named(ident.clone().unwrap()))
189            .collect::<Vec<_>>(),
190        syn::Fields::Unnamed(FieldsUnnamed { unnamed, .. }) => (0..(unnamed.len() as u32))
191            .map(|index| {
192                Member::Unnamed(Index {
193                    index,
194                    span: Span::call_site(),
195                })
196            })
197            .collect(),
198        syn::Fields::Unit => Vec::new(),
199    };
200
201    // Extend the original generics.
202    let mut combined_generics = item_struct.generics.clone();
203    combined_generics
204        .params
205        .extend(item_struct_renamed.generics.params.clone());
206
207    ProcessItemStruct {
208        root: root(),
209        item_struct,
210        item_struct_renamed,
211        self_where_predicates,
212        both_where_predicates,
213        field_names,
214        combined_generics,
215    }
216}
217
218/// See [`derive_lattice_macro`].
219fn derive_lattice(process_item_struct: &ProcessItemStruct) -> TokenStream {
220    let mut out = TokenStream::new();
221    out.extend(derive_merge(process_item_struct));
222    out.extend(derive_lattice_ord(process_item_struct));
223    out.extend(derive_is_bot(process_item_struct));
224    out.extend(derive_is_top(process_item_struct));
225    out.extend(derive_lattice_from(process_item_struct));
226    out
227}
228
229/// See [`derive_merge_macro`].
230fn derive_merge(
231    ProcessItemStruct {
232        root,
233        item_struct,
234        item_struct_renamed,
235        self_where_predicates: _,
236        both_where_predicates,
237        field_names,
238        combined_generics,
239    }: &ProcessItemStruct,
240) -> TokenStream {
241    let merge_where_predicates = item_struct
242        .fields
243        .iter()
244        .zip(item_struct_renamed.fields.iter())
245        .map(|(field_self, field_othr)| {
246            let ty_self = &field_self.ty;
247            let ty_othr = &field_othr.ty;
248            quote! {
249                #ty_self: #root::Merge<#ty_othr>
250            }
251        });
252
253    let ident = &item_struct.ident;
254    let (_, ty_generics_self, _) = item_struct.generics.split_for_impl();
255    let (_, ty_generics_othr, _) = item_struct_renamed.generics.split_for_impl();
256    let (impl_generics_both, _, _) = combined_generics.split_for_impl();
257    quote! {
258        impl #impl_generics_both #root::Merge<#ident #ty_generics_othr> for #ident #ty_generics_self
259        where
260            #both_where_predicates
261            #( #merge_where_predicates ),*
262        {
263            fn merge(&mut self, other: #ident #ty_generics_othr) -> bool {
264                let mut changed = false;
265                #(
266                    changed |= #root::Merge::merge(&mut self.#field_names, other.#field_names);
267                )*
268                changed
269            }
270        }
271    }
272}
273
274/// See [`derive_lattice_ord_macro`].
275fn derive_lattice_ord(
276    ProcessItemStruct {
277        root,
278        item_struct,
279        item_struct_renamed,
280        self_where_predicates: _,
281        both_where_predicates,
282        field_names,
283        combined_generics,
284    }: &ProcessItemStruct,
285) -> TokenStream {
286    // PartialEq.
287    let pareq_where_predicates = item_struct
288        .fields
289        .iter()
290        .zip(item_struct_renamed.fields.iter())
291        .map(|(field_self, field_othr)| {
292            let ty_self = &field_self.ty;
293            let ty_othr = &field_othr.ty;
294            quote! {
295                #ty_self: ::core::cmp::PartialEq<#ty_othr>
296            }
297        });
298    // PartialOrd and LatticeOrd.
299    let compare_where_predicates = item_struct
300        .fields
301        .iter()
302        .zip(item_struct_renamed.fields.iter())
303        .map(|(field_self, field_othr)| {
304            let ty_self = &field_self.ty;
305            let ty_othr = &field_othr.ty;
306            quote! {
307                #ty_self: ::core::cmp::PartialOrd<#ty_othr>
308            }
309        })
310        .collect::<Vec<_>>();
311
312    let ident = &item_struct.ident;
313    let (_, ty_generics_self, _) = item_struct.generics.split_for_impl();
314    let (_, ty_generics_othr, _) = item_struct_renamed.generics.split_for_impl();
315    let (impl_generics_both, _, _) = combined_generics.split_for_impl();
316    quote! {
317        impl #impl_generics_both ::core::cmp::PartialEq<#ident #ty_generics_othr> for #ident #ty_generics_self
318        where
319            #both_where_predicates
320            #( #pareq_where_predicates ),*
321        {
322            fn eq(&self, other: &#ident #ty_generics_othr) -> bool {
323                #(
324                    if !::core::cmp::PartialEq::eq(&self.#field_names, &other.#field_names) {
325                        return false;
326                    }
327                )*
328                true
329            }
330        }
331
332        impl #impl_generics_both ::core::cmp::PartialOrd<#ident #ty_generics_othr> for #ident #ty_generics_self
333        where
334            #both_where_predicates
335            #( #compare_where_predicates ),*
336        {
337            fn partial_cmp(&self, other: &#ident #ty_generics_othr) -> ::core::option::Option<::core::cmp::Ordering> {
338                let mut self_any_greater = false;
339                let mut othr_any_greater = false;
340                #(
341                    // `?` short-circuits `None` (uncomparable).
342                    match ::core::cmp::PartialOrd::partial_cmp(&self.#field_names, &other.#field_names)? {
343                        ::core::cmp::Ordering::Less => {
344                            othr_any_greater = true;
345                        }
346                        ::core::cmp::Ordering::Greater => {
347                            self_any_greater = true;
348                        }
349                        ::core::cmp::Ordering::Equal => {}
350                    }
351                    if self_any_greater && othr_any_greater {
352                        return ::core::option::Option::None;
353                    }
354                )*
355                ::core::option::Option::Some(
356                    match (self_any_greater, othr_any_greater) {
357                        (false, false) => ::core::cmp::Ordering::Equal,
358                        (false, true) => ::core::cmp::Ordering::Less,
359                        (true, false) => ::core::cmp::Ordering::Greater,
360                        (true, true) => ::core::unreachable!(),
361                    }
362                )
363            }
364        }
365        impl #impl_generics_both #root::LatticeOrd<#ident #ty_generics_othr> for #ident #ty_generics_self
366        where
367            #both_where_predicates
368            #( #compare_where_predicates ),*
369        {}
370    }
371}
372
373/// See [`derive_is_bot_macro`].
374fn derive_is_bot(
375    ProcessItemStruct {
376        root,
377        item_struct,
378        item_struct_renamed: _,
379        self_where_predicates,
380        both_where_predicates: _,
381        field_names,
382        combined_generics: _,
383    }: &ProcessItemStruct,
384) -> TokenStream {
385    let isbot_where_predicates = item_struct.fields.iter().map(|Field { ty, .. }| {
386        quote! {
387            #ty: #root::IsBot
388        }
389    });
390
391    let ident = &item_struct.ident;
392    let (impl_generics_self, ty_generics_self, _) = item_struct.generics.split_for_impl();
393    quote! {
394        impl #impl_generics_self #root::IsBot for #ident #ty_generics_self
395        where
396            #self_where_predicates
397            #( #isbot_where_predicates ),*
398        {
399            fn is_bot(&self) -> bool {
400                #(
401                    if !#root::IsBot::is_bot(&self.#field_names) {
402                        return false;
403                    }
404                )*
405                true
406            }
407        }
408    }
409}
410
411/// See [`derive_is_top_macro`].
412fn derive_is_top(
413    ProcessItemStruct {
414        root,
415        item_struct,
416        item_struct_renamed: _,
417        self_where_predicates,
418        both_where_predicates: _,
419        field_names,
420        combined_generics: _,
421    }: &ProcessItemStruct,
422) -> TokenStream {
423    let istop_where_predicates = item_struct.fields.iter().map(|Field { ty, .. }| {
424        quote! {
425            #ty: #root::IsTop
426        }
427    });
428
429    let ident = &item_struct.ident;
430    let (impl_generics_self, ty_generics_self, _) = item_struct.generics.split_for_impl();
431    quote! {
432        impl #impl_generics_self #root::IsTop for #ident #ty_generics_self
433        where
434            #self_where_predicates
435            #( #istop_where_predicates ),*
436        {
437            fn is_top(&self) -> bool {
438                #(
439                    if !#root::IsTop::is_top(&self.#field_names) {
440                        return false;
441                    }
442                )*
443                true
444            }
445        }
446    }
447}
448
449/// See [`derive_lattice_from_macro`].
450fn derive_lattice_from(
451    ProcessItemStruct {
452        root,
453        item_struct,
454        item_struct_renamed,
455        self_where_predicates: _,
456        both_where_predicates,
457        field_names,
458        combined_generics,
459    }: &ProcessItemStruct,
460) -> TokenStream {
461    let latticefrom_where_predicates = item_struct
462        .fields
463        .iter()
464        .zip(item_struct_renamed.fields.iter())
465        .map(|(field_self, field_othr)| {
466            let ty_self = &field_self.ty;
467            let ty_othr = &field_othr.ty;
468            quote! {
469                #ty_self: #root::LatticeFrom<#ty_othr>
470            }
471        });
472
473    let ident = &item_struct.ident;
474    let (_, ty_generics_self, _) = item_struct.generics.split_for_impl();
475    let (_, ty_generics_othr, _) = item_struct_renamed.generics.split_for_impl();
476    let (impl_generics_both, _, _) = combined_generics.split_for_impl();
477    quote! {
478        impl #impl_generics_both #root::LatticeFrom<#ident #ty_generics_othr> for #ident #ty_generics_self
479        where
480            #both_where_predicates
481            #( #latticefrom_where_predicates ),*
482        {
483            fn lattice_from(other: #ident #ty_generics_othr) -> Self {
484                Self {
485                    #(
486                        #field_names: #root::LatticeFrom::lattice_from(other.#field_names),
487                    )*
488                }
489            }
490        }
491    }
492}
493
494/// Also see `lattices/tests/macro.rs`
495#[cfg(test)]
496mod test {
497    use syn::parse_quote;
498
499    use super::*;
500
501    /// Snapshots the macro output without actually testing if it compiles.
502    /// See `lattices/tests/macro.rs` for compiling tests.
503    macro_rules! assert_derive_snapshots {
504        ( $( $t:tt )* ) => {
505            {
506                let item = parse_quote! {
507                    $( $t )*
508                };
509                let process_item_struct = process_item_struct(item);
510                let derive_lattice = derive_lattice(&process_item_struct);
511                hydro_build_utils::assert_snapshot!(prettyplease::unparse(&parse_quote! { #derive_lattice }));
512            }
513        };
514    }
515
516    #[test]
517    fn derive_example() {
518        assert_derive_snapshots! {
519            struct MyLattice<KeySet, Epoch> {
520                keys: SetUnion<KeySet>,
521                epoch: Max<Epoch>,
522            }
523        };
524    }
525
526    #[test]
527    fn derive_pair() {
528        assert_derive_snapshots! {
529            pub struct Pair<LatA, LatB> {
530                pub a: LatA,
531                pub b: LatB,
532            }
533        };
534    }
535
536    #[test]
537    fn derive_similar_fields() {
538        // Will create duplicate where clauses, but that is OK.
539        assert_derive_snapshots! {
540            pub struct SimilarFields {
541                a: Max<usize>,
542                b: Max<usize>,
543                c: Max<usize>,
544            }
545        };
546    }
547}