openvm_codec_derive/
lib.rs1use proc_macro::TokenStream;
2use quote::{format_ident, quote};
3use syn::{parse_macro_input, Data, DeriveInput, Fields};
4#[cfg(feature = "lean")]
5use syn::{GenericParam, Type};
6
7fn codec_crate_root() -> proc_macro2::TokenStream {
8 match proc_macro_crate::crate_name("openvm-stark-backend") {
9 Ok(proc_macro_crate::FoundCrate::Itself) => quote!(crate),
10 Ok(proc_macro_crate::FoundCrate::Name(name)) => {
11 let ident = format_ident!("{}", name);
12 quote!(::#ident)
13 }
14 Err(_) => quote!(::openvm_stark_backend),
15 }
16}
17
18#[proc_macro_derive(Encode)]
19pub fn encode_derive(input: TokenStream) -> TokenStream {
20 let ast = parse_macro_input!(input as DeriveInput);
21 let name = &ast.ident;
22 let (impl_generics, type_generics, where_clause) = ast.generics.split_for_impl();
23 let codec_root = codec_crate_root();
24
25 let fields = match &ast.data {
26 Data::Struct(data_struct) => &data_struct.fields,
27 Data::Enum(_) => {
28 return syn::Error::new(
29 name.span(),
30 "Encode derive macro only supports structs, not enums",
31 )
32 .to_compile_error()
33 .into();
34 }
35 Data::Union(_) => {
36 return syn::Error::new(
37 name.span(),
38 "Encode derive macro only supports structs, not unions",
39 )
40 .to_compile_error()
41 .into();
42 }
43 };
44
45 let encode_fields = match fields {
46 Fields::Named(fields_named) => {
47 let field_encodes = fields_named.named.iter().map(|field| {
48 let field_name = &field.ident;
49 quote! {
50 self.#field_name.encode(writer)?;
51 }
52 });
53 quote! {
54 #(#field_encodes)*
55 }
56 }
57 Fields::Unnamed(fields_unnamed) => {
58 let field_encodes = fields_unnamed.unnamed.iter().enumerate().map(|(idx, _)| {
59 let index = syn::Index::from(idx);
60 quote! {
61 self.#index.encode(writer)?;
62 }
63 });
64 quote! {
65 #(#field_encodes)*
66 }
67 }
68 Fields::Unit => {
69 quote! {}
70 }
71 };
72
73 let expanded = quote! {
74 impl #impl_generics #codec_root::codec::Encode for #name #type_generics #where_clause {
75 fn encode<W: std::io::Write>(&self, writer: &mut W) -> std::io::Result<()> {
76 #encode_fields
77 Ok(())
78 }
79 }
80 };
81
82 TokenStream::from(expanded)
83}
84
85#[proc_macro_derive(Decode)]
86pub fn decode_derive(input: TokenStream) -> TokenStream {
87 let ast = parse_macro_input!(input as DeriveInput);
88 let name = &ast.ident;
89 let (impl_generics, type_generics, where_clause) = ast.generics.split_for_impl();
90 let codec_root = codec_crate_root();
91
92 let fields = match &ast.data {
93 Data::Struct(data_struct) => &data_struct.fields,
94 Data::Enum(_) => {
95 return syn::Error::new(
96 name.span(),
97 "Decode derive macro only supports structs, not enums",
98 )
99 .to_compile_error()
100 .into();
101 }
102 Data::Union(_) => {
103 return syn::Error::new(
104 name.span(),
105 "Decode derive macro only supports structs, not unions",
106 )
107 .to_compile_error()
108 .into();
109 }
110 };
111
112 let decode_fields = match fields {
113 Fields::Named(fields_named) => {
114 let field_decodes = fields_named.named.iter().map(|field| {
115 let field_name = &field.ident;
116 let field_ty = &field.ty;
117 quote! {
118 #field_name: <#field_ty as #codec_root::codec::Decode>::decode(reader)?,
119 }
120 });
121 quote! {
122 {
123 #(#field_decodes)*
124 }
125 }
126 }
127 Fields::Unnamed(fields_unnamed) => {
128 let field_decodes = fields_unnamed.unnamed.iter().map(|field| {
129 let field_ty = &field.ty;
130 quote! {
131 <#field_ty as #codec_root::codec::Decode>::decode(reader)?,
132 }
133 });
134 quote! {
135 (
136 #(#field_decodes)*
137 )
138 }
139 }
140 Fields::Unit => {
141 quote! {}
142 }
143 };
144
145 let expanded = quote! {
146 impl #impl_generics #codec_root::codec::Decode for #name #type_generics #where_clause {
147 fn decode<R: std::io::Read>(reader: &mut R) -> std::io::Result<Self> {
148 Ok(Self #decode_fields)
149 }
150 }
151 };
152
153 TokenStream::from(expanded)
154}
155
156#[cfg(feature = "lean")]
164#[proc_macro_derive(LeanColumns)]
165pub fn lean_columns_derive(input: TokenStream) -> TokenStream {
166 let ast = parse_macro_input!(input as DeriveInput);
167 let name = &ast.ident;
168 let (impl_generics, type_generics, where_clause) = ast.generics.split_for_impl();
169 let lean_root = codec_crate_root();
170
171 let type_generic = ast
173 .generics
174 .params
175 .iter()
176 .find_map(|param| match param {
177 GenericParam::Type(tp) => Some(&tp.ident),
178 _ => None,
179 })
180 .expect("LeanColumns requires at least one type generic");
181
182 let fields = match &ast.data {
183 Data::Struct(data_struct) => match &data_struct.fields {
184 Fields::Named(f) => &f.named,
185 _ => {
186 return syn::Error::new(name.span(), "LeanColumns only supports named fields")
187 .to_compile_error()
188 .into();
189 }
190 },
191 _ => {
192 return syn::Error::new(name.span(), "LeanColumns only supports structs")
193 .to_compile_error()
194 .into();
195 }
196 };
197
198 let push_stmts = fields.iter().map(|field| {
199 let field_name = field.ident.as_ref().unwrap().to_string();
200 let field_ty = &field.ty;
201
202 if is_type_generic(field_ty, type_generic) {
203 quote! {
205 entries.push(#lean_root::lean::LeanEntry::Column(#field_name.to_string()));
206 }
207 } else if let Type::Array(arr) = field_ty {
208 if is_type_generic(&arr.elem, type_generic) {
209 let len = &arr.len;
211 quote! {
212 for i in 0..#len {
213 entries.push(#lean_root::lean::LeanEntry::Column(
214 format!("{}_{}", #field_name, i),
215 ));
216 }
217 }
218 } else if let Type::Array(inner_arr) = &*arr.elem {
219 if is_type_generic(&inner_arr.elem, type_generic) {
220 let outer_len = &arr.len;
222 let inner_len = &inner_arr.len;
223 quote! {
224 for i in 0..#outer_len {
225 for j in 0..#inner_len {
226 entries.push(#lean_root::lean::LeanEntry::Column(
227 format!("{}_{}_{}", #field_name, i, j),
228 ));
229 }
230 }
231 }
232 } else {
233 let elem_ty = &inner_arr.elem;
235 let type_name_str = extract_type_name(elem_ty);
236 let outer_len = &arr.len;
237 let inner_len = &inner_arr.len;
238 quote! {
239 for _ in 0..#outer_len {
240 for _ in 0..#inner_len {
241 entries.push(#lean_root::lean::LeanEntry::SubAir {
242 field_name: #field_name.to_string(),
243 type_name: #type_name_str.to_string(),
244 width: std::mem::size_of::<#elem_ty>() / std::mem::size_of::<#type_generic>(),
245 });
246 }
247 }
248 }
249 }
250 } else {
251 let elem_ty = &arr.elem;
253 let type_name_str = extract_type_name(elem_ty);
254 let len = &arr.len;
255 quote! {
256 for _ in 0..#len {
257 entries.push(#lean_root::lean::LeanEntry::SubAir {
258 field_name: #field_name.to_string(),
259 type_name: #type_name_str.to_string(),
260 width: std::mem::size_of::<#elem_ty>() / std::mem::size_of::<#type_generic>(),
261 });
262 }
263 }
264 }
265 } else {
266 let type_name_str = extract_type_name(field_ty);
268 quote! {
269 entries.push(#lean_root::lean::LeanEntry::SubAir {
270 field_name: #field_name.to_string(),
271 type_name: #type_name_str.to_string(),
272 width: std::mem::size_of::<#field_ty>() / std::mem::size_of::<#type_generic>(),
273 });
274 }
275 }
276 });
277
278 let expanded = quote! {
279 impl #impl_generics #lean_root::lean::LeanColumns for #name #type_generics #where_clause {
280 fn lean_columns() -> Vec<#lean_root::lean::LeanEntry> {
281 let mut entries = Vec::new();
282 #(#push_stmts)*
283 entries
284 }
285 }
286 };
287
288 TokenStream::from(expanded)
289}
290
291#[cfg(feature = "lean")]
293fn is_type_generic(ty: &Type, generic: &syn::Ident) -> bool {
294 if let Type::Path(tp) = ty {
295 tp.qself.is_none() && tp.path.is_ident(generic)
296 } else {
297 false
298 }
299}
300
301#[cfg(feature = "lean")]
303fn extract_type_name(ty: &Type) -> String {
304 if let Type::Path(tp) = ty {
305 if let Some(segment) = tp.path.segments.last() {
306 return segment.ident.to_string();
307 }
308 }
309 "Unknown".to_string()
310}