openvm_circuit_primitives_derive/
lib.rs

1// AlignedBorrow is copied from valida-derive under MIT license
2extern crate alloc;
3extern crate proc_macro;
4
5use itertools::multiunzip;
6use proc_macro::TokenStream;
7use quote::quote;
8use syn::{parse_macro_input, Data, DeriveInput, Fields, GenericParam, LitStr, Meta};
9
10mod cols_ref;
11use cols_ref::cols_ref_impl;
12
13/// Derives `ColumnsAir` for an AIR struct.
14///
15/// `#[columns_via(SomeCols<u8, ...>)]` selects the column struct whose
16/// `StructReflectionHelper::struct_reflection()` (typically derived via
17/// `#[derive(StructReflection)]`) provides the column names. The reflection is
18/// element-type invariant, so any concrete element type works — `u8` is just
19/// the conventional choice. If the AIR has no columns to reflect, write
20/// `impl ColumnsAir for X {}` by hand instead.
21#[proc_macro_derive(ColumnsAir, attributes(columns_via))]
22pub fn columns_air_derive(input: TokenStream) -> TokenStream {
23    let ast: DeriveInput = parse_macro_input!(input as DeriveInput);
24    let name = &ast.ident;
25
26    let columns_via = match ast
27        .attrs
28        .iter()
29        .find(|attr| attr.path().is_ident("columns_via"))
30    {
31        Some(attr) => attr,
32        None => {
33            return syn::Error::new_spanned(
34                &ast.ident,
35                "#[derive(ColumnsAir)] requires a `#[columns_via(ColsTy<u8, ...>)]` attribute. \
36                 If the AIR has no columns to reflect, write `impl ColumnsAir for X {}` by hand.",
37            )
38            .to_compile_error()
39            .into();
40        }
41    };
42
43    let cols_ty: syn::Type = match columns_via.parse_args() {
44        Ok(ty) => ty,
45        Err(err) => return err.to_compile_error().into(),
46    };
47
48    let (impl_generics, ty_generics, where_clause) = ast.generics.split_for_impl();
49
50    quote! {
51        impl #impl_generics ::openvm_circuit_primitives::ColumnsAir
52            for #name #ty_generics
53            #where_clause
54        {
55            fn columns(&self) -> Option<Vec<String>> {
56                <#cols_ty as ::openvm_circuit_primitives::StructReflectionHelper>::struct_reflection()
57            }
58        }
59    }
60    .into()
61}
62
63#[proc_macro_derive(AlignedBorrow)]
64pub fn aligned_borrow_derive(input: TokenStream) -> TokenStream {
65    let ast = parse_macro_input!(input as DeriveInput);
66    let name = &ast.ident;
67
68    // Get first generic which must be type (ex. `T`) for input <T, N: NumLimbs, const M: usize>
69    let type_generic = ast
70        .generics
71        .params
72        .iter()
73        .map(|param| match param {
74            GenericParam::Type(type_param) => &type_param.ident,
75            _ => panic!("Expected first generic to be a type"),
76        })
77        .next()
78        .expect("Expected at least one generic");
79
80    // Get generics after the first (ex. `N: NumLimbs, const M: usize`)
81    // We need this because when we assert the size, we want to substitute u8 for T.
82    let non_first_generics = ast
83        .generics
84        .params
85        .iter()
86        .skip(1)
87        .filter_map(|param| match param {
88            GenericParam::Type(type_param) => Some(&type_param.ident),
89            GenericParam::Const(const_param) => Some(&const_param.ident),
90            _ => None,
91        })
92        .collect::<Vec<_>>();
93
94    // Get impl generics (`<T, N: NumLimbs, const M: usize>`), type generics (`<T, N>`), where
95    // clause (`where T: Clone`)
96    let (impl_generics, type_generics, where_clause) = ast.generics.split_for_impl();
97
98    let methods = quote! {
99        impl #impl_generics core::borrow::Borrow<#name #type_generics> for [#type_generic] #where_clause {
100            fn borrow(&self) -> &#name #type_generics {
101                debug_assert_eq!(self.len(), #name::#type_generics::width());
102                let (prefix, shorts, _suffix) = unsafe { self.align_to::<#name #type_generics>() };
103                debug_assert!(prefix.is_empty(), "Alignment should match");
104                debug_assert_eq!(shorts.len(), 1);
105                &shorts[0]
106            }
107        }
108
109        impl #impl_generics core::borrow::BorrowMut<#name #type_generics> for [#type_generic] #where_clause {
110            fn borrow_mut(&mut self) -> &mut #name #type_generics {
111                debug_assert_eq!(self.len(), #name::#type_generics::width());
112                let (prefix, shorts, _suffix) = unsafe { self.align_to_mut::<#name #type_generics>() };
113                debug_assert!(prefix.is_empty(), "Alignment should match");
114                debug_assert_eq!(shorts.len(), 1);
115                &mut shorts[0]
116            }
117        }
118
119        impl #impl_generics #name #type_generics {
120            pub const fn width() -> usize {
121                std::mem::size_of::<#name<u8 #(, #non_first_generics)*>>()
122            }
123        }
124    };
125
126    TokenStream::from(methods)
127}
128
129/// `S` is the type the derive macro is being called on
130/// Implements `Borrow<S>` and `BorrowMut<S>` for `[u8]`
131/// [u8] has to have (checked via `debug_assert!`s)
132/// - at least size_of(S) length
133/// - at least align_of(S) alignment
134#[proc_macro_derive(AlignedBytesBorrow)]
135pub fn aligned_bytes_borrow_derive(input: TokenStream) -> TokenStream {
136    let ast = parse_macro_input!(input as DeriveInput);
137    let name = &ast.ident;
138
139    // Get impl generics, type generics, where clause
140    // Note, need to add the new type generic to the `impl_generics`
141    let (impl_generics, type_generics, where_clause) = ast.generics.split_for_impl();
142
143    let methods = quote! {
144        impl #impl_generics core::borrow::Borrow<#name #type_generics> for [u8]
145        where
146            #where_clause
147        {
148            fn borrow(&self) -> &#name #type_generics {
149                use core::mem::{align_of, size_of_val};
150                debug_assert!(size_of_val(self) >= core::mem::size_of::<#name #type_generics>());
151                debug_assert_eq!(self.as_ptr() as usize % align_of::<#name #type_generics>(), 0);
152                unsafe { &*(self.as_ptr() as *const #name #type_generics) }
153            }
154        }
155
156        impl #impl_generics core::borrow::BorrowMut<#name #type_generics> for [u8]
157        where
158            #where_clause
159        {
160            fn borrow_mut(&mut self) -> &mut #name #type_generics {
161                use core::mem::{align_of, size_of_val};
162                debug_assert!(size_of_val(self) >= core::mem::size_of::<#name #type_generics>());
163                debug_assert_eq!(self.as_ptr() as usize % align_of::<#name #type_generics>(), 0);
164                unsafe { &mut *(self.as_mut_ptr() as *mut #name #type_generics) }
165            }
166        }
167    };
168
169    TokenStream::from(methods)
170}
171
172#[proc_macro_derive(Chip, attributes(chip))]
173pub fn chip_derive(input: TokenStream) -> TokenStream {
174    // Parse the attributes from the struct or enum
175    let ast: syn::DeriveInput = syn::parse(input).unwrap();
176
177    let name = &ast.ident;
178    let generics = &ast.generics;
179    let (_impl_generics, ty_generics, _where_clause) = generics.split_for_impl();
180
181    match &ast.data {
182        Data::Struct(inner) => {
183            let generics = &ast.generics;
184            let mut new_generics = generics.clone();
185            new_generics.params.push(syn::parse_quote! { R });
186            new_generics
187                .params
188                .push(syn::parse_quote! { PB: openvm_stark_backend::prover::ProverBackend });
189            let (impl_generics, _, _) = new_generics.split_for_impl();
190
191            // Check if the struct has only one unnamed field
192            let inner_ty = match &inner.fields {
193                Fields::Unnamed(fields) => {
194                    if fields.unnamed.len() != 1 {
195                        panic!("Only one unnamed field is supported");
196                    }
197                    fields.unnamed.first().unwrap().ty.clone()
198                }
199                _ => panic!("Only unnamed fields are supported"),
200            };
201            let mut new_generics = generics.clone();
202            let where_clause = new_generics.make_where_clause();
203            where_clause
204                .predicates
205                .push(syn::parse_quote! { #inner_ty: openvm_circuit::primitives::Chip<R, PB> });
206            quote! {
207                impl #impl_generics openvm_circuit::primitives::Chip<R, PB> for #name #ty_generics #where_clause {
208                    fn generate_proving_ctx(&self, records: R) -> openvm_stark_backend::prover::AirProvingContext<PB> {
209                        self.0.generate_proving_ctx(records)
210                    }
211                }
212            }.into()
213        }
214        Data::Enum(e) => {
215            let variants = e
216                .variants
217                .iter()
218                .map(|variant| {
219                    let variant_name = &variant.ident;
220
221                    let mut fields = variant.fields.iter();
222                    let field = fields.next().unwrap();
223                    assert!(fields.next().is_none(), "Only one field is supported");
224                    (variant_name, field)
225                })
226                .collect::<Vec<_>>();
227
228            let (generate_proving_ctx_arms, where_predicates): (Vec<_>, Vec<_>) =
229                variants.iter().map(|(variant_name, field)| {
230                let field_ty = &field.ty;
231                let generate_proving_ctx_arm = quote! {
232                    #name::#variant_name(x) => <#field_ty as openvm_circuit::primitives::Chip<R, PB>>::generate_proving_ctx(x, records)
233                };
234                let where_predicate =
235                    syn::parse_quote! { #field_ty: openvm_circuit::primitives::Chip<R, PB> };
236                (generate_proving_ctx_arm, where_predicate)
237            }).collect();
238
239            // Attach extra generics R and PB to the impl_generics
240            let generics = &ast.generics;
241            let mut new_generics = generics.clone();
242            new_generics.params.push(syn::parse_quote! { R });
243            new_generics
244                .params
245                .push(syn::parse_quote! { PB: openvm_stark_backend::prover::ProverBackend });
246            let (impl_generics, _, _) = new_generics.split_for_impl();
247
248            // Implement Chip whenever the inner type implements Chip
249            let mut new_generics = generics.clone();
250            let where_clause = new_generics.make_where_clause();
251            for predicate in where_predicates {
252                where_clause.predicates.push(predicate);
253            }
254            let attributes = ast.attrs.iter().find(|&attr| attr.path().is_ident("chip"));
255            if let Some(attr) = attributes {
256                let mut fail_flag = false;
257
258                match &attr.meta {
259                    Meta::List(meta_list) => {
260                        meta_list
261                            .parse_nested_meta(|meta| {
262                                if meta.path.is_ident("where") {
263                                    let value = meta.value()?; // this parses the `=`
264                                    let s: LitStr = value.parse()?;
265                                    let where_value = s.value();
266                                    where_clause.predicates.push(syn::parse_str(&where_value)?);
267                                } else {
268                                    fail_flag = true;
269                                }
270                                Ok(())
271                            })
272                            .unwrap();
273                    }
274                    _ => fail_flag = true,
275                }
276                if fail_flag {
277                    return syn::Error::new(
278                        name.span(),
279                        "Only `#[chip(where = ...)]` format is supported",
280                    )
281                    .to_compile_error()
282                    .into();
283                }
284            }
285
286            quote! {
287                impl #impl_generics openvm_circuit::primitives::Chip<R, PB> for #name #ty_generics #where_clause {
288                    fn generate_proving_ctx(&self, records: R) -> openvm_stark_backend::prover::AirProvingContext<PB> {
289                        match self {
290                            #(#generate_proving_ctx_arms,)*
291                        }
292                    }
293                }
294            }.into()
295        }
296        Data::Union(_) => unimplemented!("Unions are not supported"),
297    }
298}
299#[proc_macro_derive(BytesStateful)]
300pub fn bytes_stateful_derive(input: TokenStream) -> TokenStream {
301    let ast: syn::DeriveInput = syn::parse(input).unwrap();
302
303    let name = &ast.ident;
304    let generics = &ast.generics;
305    let (impl_generics, ty_generics, _) = generics.split_for_impl();
306
307    match &ast.data {
308        Data::Struct(inner) => {
309            // Check if the struct has only one unnamed field
310            let inner_ty = match &inner.fields {
311                Fields::Unnamed(fields) => {
312                    if fields.unnamed.len() != 1 {
313                        panic!("Only one unnamed field is supported");
314                    }
315                    fields.unnamed.first().unwrap().ty.clone()
316                }
317                _ => panic!("Only unnamed fields are supported"),
318            };
319            // Use full path ::openvm_circuit... so it can be used either within or outside the vm
320            // crate. Assume F is already generic of the field.
321            let mut new_generics = generics.clone();
322            let where_clause = new_generics.make_where_clause();
323            where_clause
324                .predicates
325                .push(syn::parse_quote! { #inner_ty: ::openvm_stark_backend::Stateful<Vec<u8>> });
326
327            quote! {
328                impl #impl_generics ::openvm_stark_backend::Stateful<Vec<u8>> for #name #ty_generics #where_clause {
329                    fn load_state(&mut self, state: Vec<u8>) {
330                        self.0.load_state(state)
331                    }
332
333                    fn store_state(&self) -> Vec<u8> {
334                        self.0.store_state()
335                    }
336                }
337            }
338            .into()
339        }
340        Data::Enum(e) => {
341            let variants = e
342                .variants
343                .iter()
344                .map(|variant| {
345                    let variant_name = &variant.ident;
346
347                    let mut fields = variant.fields.iter();
348                    let field = fields.next().unwrap();
349                    assert!(fields.next().is_none(), "Only one field is supported");
350                    (variant_name, field)
351                })
352                .collect::<Vec<_>>();
353            // Use full path ::openvm_stark_backend... so it can be used either within or outside
354            // the vm crate.
355            let (load_state_arms, store_state_arms): (Vec<_>, Vec<_>) =
356                multiunzip(variants.iter().map(|(variant_name, field)| {
357                    let field_ty = &field.ty;
358                    let load_state_arm = quote! {
359                        #name::#variant_name(x) => <#field_ty as ::openvm_stark_backend::Stateful<Vec<u8>>>::load_state(x, state)
360                    };
361                    let store_state_arm = quote! {
362                        #name::#variant_name(x) => <#field_ty as ::openvm_stark_backend::Stateful<Vec<u8>>>::store_state(x)
363                    };
364
365                    (load_state_arm, store_state_arm)
366                }));
367            quote! {
368                impl #impl_generics ::openvm_stark_backend::Stateful<Vec<u8>> for #name #ty_generics {
369                    fn load_state(&mut self, state: Vec<u8>) {
370                        match self {
371                            #(#load_state_arms,)*
372                        }
373                    }
374
375                    fn store_state(&self) -> Vec<u8> {
376                        match self {
377                            #(#store_state_arms,)*
378                        }
379                    }
380                }
381            }
382            .into()
383        }
384        _ => unimplemented!(),
385    }
386}
387
388#[proc_macro_derive(ColsRef, attributes(aligned_borrow, config))]
389pub fn cols_ref_derive(input: TokenStream) -> TokenStream {
390    let derive_input: DeriveInput = parse_macro_input!(input as DeriveInput);
391
392    let config = derive_input
393        .attrs
394        .iter()
395        .find(|attr| attr.path().is_ident("config"));
396    if config.is_none() {
397        return syn::Error::new(derive_input.ident.span(), "Config attribute is required")
398            .to_compile_error()
399            .into();
400    }
401    let config: proc_macro2::Ident = config
402        .unwrap()
403        .parse_args()
404        .expect("Failed to parse config");
405
406    let res = cols_ref_impl(derive_input, config);
407    res.into()
408}