Skip to main content

openvm_codec_derive/
lib.rs

1use 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/// Generates `fn lean_columns() -> Vec<LeanEntry>` for Cols structs.
157///
158/// - `field: T` → `LeanEntry::Column("field")`
159/// - `field: [T; N]` → `LeanEntry::Column("field_0")` .. `LeanEntry::Column("field_{N-1}")`
160/// - `field: [[T; N]; M]` → `LeanEntry::Column("field_0_0")` ..
161///   `LeanEntry::Column("field_{M-1}_{N-1}")`
162/// - `field: SomeStruct<T, ..>` → `LeanEntry::SubAir { field_name, type_name, width }`
163#[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    // Get the first type generic (e.g. `T`)
172    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            // Scalar field: T → Column("field")
204            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                // Array of T: [T; N] → Column("field_0") .. Column("field_{N-1}")
210                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                    // Nested array of T: [[T; N]; M] → Column("field_0_0") .. Column("field_{M-1}_{N-1}")
221                    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                    // Array of array of structs: [[SomeStruct<T>; N]; M] → M*N SubAir entries
234                    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                // Array of structs: [SomeStruct<T>; N] → N SubAir entries
252                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            // Nested struct: SomeStruct<T, ..> → SubAir
267            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/// Check if a type is exactly the given generic ident (e.g. `T`).
292#[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/// Extract the top-level type name from a type (e.g. `ExecutionState<T>` → `"ExecutionState"`).
302#[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}