openvm_recursion_circuit_derive/
lib.rs1use 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 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 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 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 #[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 #[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}