openvm_circuit_primitives_derive/
lib.rs1extern 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#[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 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 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 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#[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 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 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 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 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 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()?; 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 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 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 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}