1#![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
15fn 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 quote! { crate }
39 } else {
40 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
52fn 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
103fn 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#[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#[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#[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#[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#[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
154struct 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}
164fn 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 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 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 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 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
218fn 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
229fn 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
274fn 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 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 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 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
373fn 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
411fn 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
449fn 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#[cfg(test)]
496mod test {
497 use syn::parse_quote;
498
499 use super::*;
500
501 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 assert_derive_snapshots! {
540 pub struct SimilarFields {
541 a: Max<usize>,
542 b: Max<usize>,
543 c: Max<usize>,
544 }
545 };
546 }
547}