extern crate alloc;
extern crate proc_macro;
use itertools::multiunzip;
use proc_macro::TokenStream;
use quote::quote;
use syn::{parse_macro_input, Data, DeriveInput, Fields, GenericParam, LitStr, Meta};
mod cols_ref;
use cols_ref::cols_ref_impl;
#[proc_macro_derive(ColumnsAir, attributes(columns_via))]
pub fn columns_air_derive(input: TokenStream) -> TokenStream {
let ast: DeriveInput = parse_macro_input!(input as DeriveInput);
let name = &ast.ident;
let columns_via = match ast
.attrs
.iter()
.find(|attr| attr.path().is_ident("columns_via"))
{
Some(attr) => attr,
None => {
return syn::Error::new_spanned(
&ast.ident,
"#[derive(ColumnsAir)] requires a `#[columns_via(ColsTy<u8, ...>)]` attribute. \
If the AIR has no columns to reflect, write `impl ColumnsAir for X {}` by hand.",
)
.to_compile_error()
.into();
}
};
let cols_ty: syn::Type = match columns_via.parse_args() {
Ok(ty) => ty,
Err(err) => return err.to_compile_error().into(),
};
let (impl_generics, ty_generics, where_clause) = ast.generics.split_for_impl();
quote! {
impl #impl_generics ::openvm_circuit_primitives::ColumnsAir
for #name #ty_generics
#where_clause
{
fn columns(&self) -> Option<Vec<String>> {
<#cols_ty as ::openvm_circuit_primitives::StructReflectionHelper>::struct_reflection()
}
}
}
.into()
}
#[proc_macro_derive(AlignedBorrow)]
pub fn aligned_borrow_derive(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as DeriveInput);
let name = &ast.ident;
let type_generic = ast
.generics
.params
.iter()
.map(|param| match param {
GenericParam::Type(type_param) => &type_param.ident,
_ => panic!("Expected first generic to be a type"),
})
.next()
.expect("Expected at least one generic");
let non_first_generics = ast
.generics
.params
.iter()
.skip(1)
.filter_map(|param| match param {
GenericParam::Type(type_param) => Some(&type_param.ident),
GenericParam::Const(const_param) => Some(&const_param.ident),
_ => None,
})
.collect::<Vec<_>>();
let (impl_generics, type_generics, where_clause) = ast.generics.split_for_impl();
let methods = quote! {
impl #impl_generics core::borrow::Borrow<#name #type_generics> for [#type_generic] #where_clause {
fn borrow(&self) -> &#name #type_generics {
debug_assert_eq!(self.len(), #name::#type_generics::width());
let (prefix, shorts, _suffix) = unsafe { self.align_to::<#name #type_generics>() };
debug_assert!(prefix.is_empty(), "Alignment should match");
debug_assert_eq!(shorts.len(), 1);
&shorts[0]
}
}
impl #impl_generics core::borrow::BorrowMut<#name #type_generics> for [#type_generic] #where_clause {
fn borrow_mut(&mut self) -> &mut #name #type_generics {
debug_assert_eq!(self.len(), #name::#type_generics::width());
let (prefix, shorts, _suffix) = unsafe { self.align_to_mut::<#name #type_generics>() };
debug_assert!(prefix.is_empty(), "Alignment should match");
debug_assert_eq!(shorts.len(), 1);
&mut shorts[0]
}
}
impl #impl_generics #name #type_generics {
pub const fn width() -> usize {
std::mem::size_of::<#name<u8 #(, #non_first_generics)*>>()
}
}
};
TokenStream::from(methods)
}
#[proc_macro_derive(AlignedBytesBorrow)]
pub fn aligned_bytes_borrow_derive(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as DeriveInput);
let name = &ast.ident;
let (impl_generics, type_generics, where_clause) = ast.generics.split_for_impl();
let methods = quote! {
impl #impl_generics core::borrow::Borrow<#name #type_generics> for [u8]
where
#where_clause
{
fn borrow(&self) -> &#name #type_generics {
use core::mem::{align_of, size_of_val};
debug_assert!(size_of_val(self) >= core::mem::size_of::<#name #type_generics>());
debug_assert_eq!(self.as_ptr() as usize % align_of::<#name #type_generics>(), 0);
unsafe { &*(self.as_ptr() as *const #name #type_generics) }
}
}
impl #impl_generics core::borrow::BorrowMut<#name #type_generics> for [u8]
where
#where_clause
{
fn borrow_mut(&mut self) -> &mut #name #type_generics {
use core::mem::{align_of, size_of_val};
debug_assert!(size_of_val(self) >= core::mem::size_of::<#name #type_generics>());
debug_assert_eq!(self.as_ptr() as usize % align_of::<#name #type_generics>(), 0);
unsafe { &mut *(self.as_mut_ptr() as *mut #name #type_generics) }
}
}
};
TokenStream::from(methods)
}
#[proc_macro_derive(Chip, attributes(chip))]
pub fn chip_derive(input: TokenStream) -> TokenStream {
let ast: syn::DeriveInput = syn::parse(input).unwrap();
let name = &ast.ident;
let generics = &ast.generics;
let (_impl_generics, ty_generics, _where_clause) = generics.split_for_impl();
match &ast.data {
Data::Struct(inner) => {
let generics = &ast.generics;
let mut new_generics = generics.clone();
new_generics.params.push(syn::parse_quote! { R });
new_generics
.params
.push(syn::parse_quote! { PB: openvm_stark_backend::prover::ProverBackend });
let (impl_generics, _, _) = new_generics.split_for_impl();
let inner_ty = match &inner.fields {
Fields::Unnamed(fields) => {
if fields.unnamed.len() != 1 {
panic!("Only one unnamed field is supported");
}
fields.unnamed.first().unwrap().ty.clone()
}
_ => panic!("Only unnamed fields are supported"),
};
let mut new_generics = generics.clone();
let where_clause = new_generics.make_where_clause();
where_clause
.predicates
.push(syn::parse_quote! { #inner_ty: openvm_circuit::primitives::Chip<R, PB> });
quote! {
impl #impl_generics openvm_circuit::primitives::Chip<R, PB> for #name #ty_generics #where_clause {
fn generate_proving_ctx(&self, records: R) -> openvm_stark_backend::prover::AirProvingContext<PB> {
self.0.generate_proving_ctx(records)
}
}
}.into()
}
Data::Enum(e) => {
let variants = e
.variants
.iter()
.map(|variant| {
let variant_name = &variant.ident;
let mut fields = variant.fields.iter();
let field = fields.next().unwrap();
assert!(fields.next().is_none(), "Only one field is supported");
(variant_name, field)
})
.collect::<Vec<_>>();
let (generate_proving_ctx_arms, where_predicates): (Vec<_>, Vec<_>) =
variants.iter().map(|(variant_name, field)| {
let field_ty = &field.ty;
let generate_proving_ctx_arm = quote! {
#name::#variant_name(x) => <#field_ty as openvm_circuit::primitives::Chip<R, PB>>::generate_proving_ctx(x, records)
};
let where_predicate =
syn::parse_quote! { #field_ty: openvm_circuit::primitives::Chip<R, PB> };
(generate_proving_ctx_arm, where_predicate)
}).collect();
let generics = &ast.generics;
let mut new_generics = generics.clone();
new_generics.params.push(syn::parse_quote! { R });
new_generics
.params
.push(syn::parse_quote! { PB: openvm_stark_backend::prover::ProverBackend });
let (impl_generics, _, _) = new_generics.split_for_impl();
let mut new_generics = generics.clone();
let where_clause = new_generics.make_where_clause();
for predicate in where_predicates {
where_clause.predicates.push(predicate);
}
let attributes = ast.attrs.iter().find(|&attr| attr.path().is_ident("chip"));
if let Some(attr) = attributes {
let mut fail_flag = false;
match &attr.meta {
Meta::List(meta_list) => {
meta_list
.parse_nested_meta(|meta| {
if meta.path.is_ident("where") {
let value = meta.value()?; let s: LitStr = value.parse()?;
let where_value = s.value();
where_clause.predicates.push(syn::parse_str(&where_value)?);
} else {
fail_flag = true;
}
Ok(())
})
.unwrap();
}
_ => fail_flag = true,
}
if fail_flag {
return syn::Error::new(
name.span(),
"Only `#[chip(where = ...)]` format is supported",
)
.to_compile_error()
.into();
}
}
quote! {
impl #impl_generics openvm_circuit::primitives::Chip<R, PB> for #name #ty_generics #where_clause {
fn generate_proving_ctx(&self, records: R) -> openvm_stark_backend::prover::AirProvingContext<PB> {
match self {
#(#generate_proving_ctx_arms,)*
}
}
}
}.into()
}
Data::Union(_) => unimplemented!("Unions are not supported"),
}
}
#[proc_macro_derive(BytesStateful)]
pub fn bytes_stateful_derive(input: TokenStream) -> TokenStream {
let ast: syn::DeriveInput = syn::parse(input).unwrap();
let name = &ast.ident;
let generics = &ast.generics;
let (impl_generics, ty_generics, _) = generics.split_for_impl();
match &ast.data {
Data::Struct(inner) => {
let inner_ty = match &inner.fields {
Fields::Unnamed(fields) => {
if fields.unnamed.len() != 1 {
panic!("Only one unnamed field is supported");
}
fields.unnamed.first().unwrap().ty.clone()
}
_ => panic!("Only unnamed fields are supported"),
};
let mut new_generics = generics.clone();
let where_clause = new_generics.make_where_clause();
where_clause
.predicates
.push(syn::parse_quote! { #inner_ty: ::openvm_stark_backend::Stateful<Vec<u8>> });
quote! {
impl #impl_generics ::openvm_stark_backend::Stateful<Vec<u8>> for #name #ty_generics #where_clause {
fn load_state(&mut self, state: Vec<u8>) {
self.0.load_state(state)
}
fn store_state(&self) -> Vec<u8> {
self.0.store_state()
}
}
}
.into()
}
Data::Enum(e) => {
let variants = e
.variants
.iter()
.map(|variant| {
let variant_name = &variant.ident;
let mut fields = variant.fields.iter();
let field = fields.next().unwrap();
assert!(fields.next().is_none(), "Only one field is supported");
(variant_name, field)
})
.collect::<Vec<_>>();
let (load_state_arms, store_state_arms): (Vec<_>, Vec<_>) =
multiunzip(variants.iter().map(|(variant_name, field)| {
let field_ty = &field.ty;
let load_state_arm = quote! {
#name::#variant_name(x) => <#field_ty as ::openvm_stark_backend::Stateful<Vec<u8>>>::load_state(x, state)
};
let store_state_arm = quote! {
#name::#variant_name(x) => <#field_ty as ::openvm_stark_backend::Stateful<Vec<u8>>>::store_state(x)
};
(load_state_arm, store_state_arm)
}));
quote! {
impl #impl_generics ::openvm_stark_backend::Stateful<Vec<u8>> for #name #ty_generics {
fn load_state(&mut self, state: Vec<u8>) {
match self {
#(#load_state_arms,)*
}
}
fn store_state(&self) -> Vec<u8> {
match self {
#(#store_state_arms,)*
}
}
}
}
.into()
}
_ => unimplemented!(),
}
}
#[proc_macro_derive(ColsRef, attributes(aligned_borrow, config))]
pub fn cols_ref_derive(input: TokenStream) -> TokenStream {
let derive_input: DeriveInput = parse_macro_input!(input as DeriveInput);
let config = derive_input
.attrs
.iter()
.find(|attr| attr.path().is_ident("config"));
if config.is_none() {
return syn::Error::new(derive_input.ident.span(), "Config attribute is required")
.to_compile_error()
.into();
}
let config: proc_macro2::Ident = config
.unwrap()
.parse_args()
.expect("Failed to parse config");
let res = cols_ref_impl(derive_input, config);
res.into()
}