openvm_recursion_circuit_derive/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3use syn::{parse_macro_input, DeriveInput, GenericParam};
4
5#[proc_macro_derive(AlignedBorrow)]
6pub fn aligned_borrow_derive(input: TokenStream) -> TokenStream {
7    let ast = parse_macro_input!(input as DeriveInput);
8    let name = &ast.ident;
9
10    // Get first generic which must be type (ex. `T`) for input <T, N: NumLimbs, const M: usize>
11    let type_generic = ast
12        .generics
13        .params
14        .iter()
15        .map(|param| match param {
16            GenericParam::Type(type_param) => &type_param.ident,
17            _ => panic!("Expected first generic to be a type"),
18        })
19        .next()
20        .expect("Expected at least one generic");
21
22    // Get generics after the first (ex. `N: NumLimbs, const M: usize`)
23    // We need this because when we assert the size, we want to substitute u8 for T.
24    let non_first_generics = ast
25        .generics
26        .params
27        .iter()
28        .skip(1)
29        .filter_map(|param| match param {
30            GenericParam::Type(type_param) => Some(&type_param.ident),
31            GenericParam::Const(const_param) => Some(&const_param.ident),
32            _ => None,
33        })
34        .collect::<Vec<_>>();
35
36    // Get impl generics (`<T, N: NumLimbs, const M: usize>`), type generics (`<T, N>`), where
37    // clause (`where T: Clone`)
38    let (impl_generics, type_generics, where_clause) = ast.generics.split_for_impl();
39
40    let methods = quote! {
41        impl #impl_generics core::borrow::Borrow<#name #type_generics> for [#type_generic] #where_clause {
42            fn borrow(&self) -> &#name #type_generics {
43                debug_assert_eq!(self.len(), #name::#type_generics::width());
44                let (prefix, shorts, _suffix) = unsafe { self.align_to::<#name #type_generics>() };
45                debug_assert!(prefix.is_empty(), "Alignment should match");
46                debug_assert_eq!(shorts.len(), 1);
47                &shorts[0]
48            }
49        }
50
51        impl #impl_generics core::borrow::BorrowMut<#name #type_generics> for [#type_generic] #where_clause {
52            fn borrow_mut(&mut self) -> &mut #name #type_generics {
53                debug_assert_eq!(self.len(), #name::#type_generics::width());
54                let (prefix, shorts, _suffix) = unsafe { self.align_to_mut::<#name #type_generics>() };
55                debug_assert!(prefix.is_empty(), "Alignment should match");
56                debug_assert_eq!(shorts.len(), 1);
57                &mut shorts[0]
58            }
59        }
60
61        impl #impl_generics #name #type_generics {
62            pub const fn width() -> usize {
63                std::mem::size_of::<#name<u8 #(, #non_first_generics)*>>()
64            }
65
66            #[inline]
67            pub fn as_slice(&self) -> &[#type_generic] {
68                debug_assert_eq!(core::mem::align_of::<#name #type_generics>(), core::mem::align_of::<#type_generic>());
69                debug_assert_eq!(core::mem::size_of::<#name #type_generics>() % core::mem::size_of::<#type_generic>(), 0);
70                unsafe {
71                    core::slice::from_raw_parts(
72                        (self as *const Self).cast::<#type_generic>(),
73                        Self::width(),
74                    )
75                }
76            }
77
78            /// Mutable view of the whole struct as a contiguous slice of T.
79             #[inline]
80             pub fn as_slice_mut(&mut self) -> &mut [#type_generic] {
81                 debug_assert_eq!(core::mem::align_of::<#name #type_generics>(), core::mem::align_of::<#type_generic>());
82                 debug_assert_eq!(core::mem::size_of::<#name #type_generics>() % core::mem::size_of::<#type_generic>(), 0);
83                 unsafe {
84                     core::slice::from_raw_parts_mut(
85                         (self as *mut Self).cast::<#type_generic>(),
86                         Self::width(),
87                     )
88                 }
89             }
90
91             /// Copy out the contents as a Vec<T>.
92             /// Requires T: Clone for the slice `.to_vec()`.
93             #[inline]
94             pub fn to_vec(&self) -> std::vec::Vec<#type_generic>
95             where
96                 #type_generic: Clone
97             {
98                 self.as_slice().to_vec()
99             }
100        }
101    };
102
103    TokenStream::from(methods)
104}