1#![deny(warnings)]
118
119extern crate proc_macro;
120
121use proc_macro::TokenStream;
122use proc_macro2::TokenStream as TokenStream2;
123use quote::{format_ident, quote};
124use syn::{
125 Data, DeriveInput, Error, Field, Fields, Index, Variant, parse_macro_input,
126 punctuated::Punctuated, spanned::Spanned, token::Comma,
127};
128
129enum StructFields<'a> {
131 Named(&'a Punctuated<Field, Comma>),
132 Unnamed(&'a Punctuated<Field, Comma>),
133}
134
135fn extract_struct_fields<'a>(
137 input: &'a DeriveInput,
138 trait_name: &str,
139) -> Result<StructFields<'a>, Error> {
140 let name = &input.ident;
141 match &input.data {
142 Data::Struct(data) => match &data.fields {
143 Fields::Named(fields) => Ok(StructFields::Named(&fields.named)),
144 Fields::Unnamed(fields) => Ok(StructFields::Unnamed(&fields.unnamed)),
145 Fields::Unit => Err(Error::new(
146 input.span(),
147 format!("{trait_name} cannot be derived for unit struct `{name}`"),
148 )),
149 },
150 Data::Enum(_) => Err(Error::new(input.span(), enum_mismatch_msg(trait_name, name))),
151 Data::Union(_) => Err(Error::new(
152 input.span(),
153 format!("{trait_name} cannot be derived for union `{name}`"),
154 )),
155 }
156}
157
158fn extract_enum_variants<'a>(
160 input: &'a DeriveInput,
161 trait_name: &str,
162) -> Result<&'a Punctuated<Variant, Comma>, Error> {
163 let name = &input.ident;
164 match &input.data {
165 Data::Enum(data) => Ok(&data.variants),
166 Data::Struct(_) => Err(Error::new(input.span(), struct_mismatch_msg(trait_name, name))),
167 Data::Union(_) => Err(Error::new(
168 input.span(),
169 format!("{trait_name} cannot be derived for union `{name}`"),
170 )),
171 }
172}
173
174fn struct_mismatch_msg(trait_name: &str, name: &syn::Ident) -> String {
175 format!("{trait_name} cannot be derived for struct `{name}`")
176}
177
178fn enum_mismatch_msg(trait_name: &str, name: &syn::Ident) -> String {
179 format!("{trait_name} cannot be derived for enum `{name}`")
180}
181
182fn ensure_no_explicit_discriminants(
184 variants: &Punctuated<Variant, Comma>,
185 trait_name: &str,
186 enum_name: &syn::Ident,
187) -> Result<(), Error> {
188 for variant in variants {
189 if variant.discriminant.is_some() {
190 return Err(Error::new(
191 variant.span(),
192 format!(
193 "{trait_name} cannot be derived for enum `{enum_name}` with explicit \
194 discriminants"
195 ),
196 ));
197 }
198 }
199 Ok(())
200}
201
202#[proc_macro_derive(DeriveFromFeltRepr)]
221pub fn derive_from_felt_repr(input: TokenStream) -> TokenStream {
222 let input = parse_macro_input!(input as DeriveInput);
223
224 let expanded = derive_from_felt_repr_impl(
225 &input,
226 quote!(miden_field_repr),
227 quote!(miden_field_repr::Felt),
228 );
229 match expanded {
230 Ok(ts) => ts,
231 Err(err) => err.into_compile_error().into(),
232 }
233}
234
235fn derive_from_felt_repr_impl(
236 input: &DeriveInput,
237 felt_repr_crate: TokenStream2,
238 felt_ty: TokenStream2,
239) -> Result<TokenStream, Error> {
240 let name = &input.ident;
241 let generics = &input.generics;
242 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
243
244 let trait_name = "FromFeltRepr";
245 let expanded = match &input.data {
246 Data::Struct(_) => match extract_struct_fields(input, trait_name)? {
247 StructFields::Named(fields) => {
248 let field_names: Vec<_> =
249 fields.iter().map(|field| field.ident.as_ref().unwrap()).collect();
250 let field_types: Vec<_> = fields.iter().map(|field| &field.ty).collect();
251 quote! {
252 impl #impl_generics #felt_repr_crate::FromFeltRepr for #name #ty_generics #where_clause {
253 #[inline(always)]
254 fn from_felt_repr(reader: &mut #felt_repr_crate::FeltReader<'_>) -> #felt_repr_crate::FeltReprResult<Self> {
255 Ok(Self {
256 #(#field_names: <#field_types as #felt_repr_crate::FromFeltRepr>::from_felt_repr(reader)?),*
257 })
258 }
259 }
260 }
261 }
262 StructFields::Unnamed(fields) => {
263 let field_types: Vec<_> = fields.iter().map(|field| &field.ty).collect();
264 let reads = field_types.iter().map(|ty| {
265 quote! { <#ty as #felt_repr_crate::FromFeltRepr>::from_felt_repr(reader)? }
266 });
267 quote! {
268 impl #impl_generics #felt_repr_crate::FromFeltRepr for #name #ty_generics #where_clause {
269 #[inline(always)]
270 fn from_felt_repr(reader: &mut #felt_repr_crate::FeltReader<'_>) -> #felt_repr_crate::FeltReprResult<Self> {
271 Ok(Self(#(#reads),*))
272 }
273 }
274 }
275 }
276 },
277 Data::Enum(_) => {
278 let variants = extract_enum_variants(input, trait_name)?;
279 ensure_no_explicit_discriminants(variants, trait_name, name)?;
280
281 let arms = variants.iter().enumerate().map(|(variant_ordinal, variant)| {
282 let variant_ident = &variant.ident;
283 let tag = variant_ordinal as u32;
284 match &variant.fields {
285 Fields::Unit => quote! { #tag => Ok(Self::#variant_ident) },
286 Fields::Unnamed(fields) => {
287 let field_types: Vec<_> = fields.unnamed.iter().map(|f| &f.ty).collect();
288 let reads = field_types.iter().map(|ty| {
289 quote! { <#ty as #felt_repr_crate::FromFeltRepr>::from_felt_repr(reader)? }
290 });
291 quote! { #tag => Ok(Self::#variant_ident(#(#reads),*)) }
292 }
293 Fields::Named(fields) => {
294 let field_idents: Vec<_> = fields
295 .named
296 .iter()
297 .map(|f| f.ident.as_ref().expect("named field"))
298 .collect();
299 let field_types: Vec<_> = fields.named.iter().map(|f| &f.ty).collect();
300 let reads = field_idents.iter().zip(field_types.iter()).map(|(ident, ty)| {
301 quote! { #ident: <#ty as #felt_repr_crate::FromFeltRepr>::from_felt_repr(reader)? }
302 });
303 quote! { #tag => Ok(Self::#variant_ident { #(#reads),* }) }
304 }
305 }
306 });
307
308 quote! {
309 impl #impl_generics #felt_repr_crate::FromFeltRepr for #name #ty_generics #where_clause {
310 #[inline(always)]
311 fn from_felt_repr(reader: &mut #felt_repr_crate::FeltReader<'_>) -> #felt_repr_crate::FeltReprResult<Self> {
312 let tag_pos = reader.pos();
313 let len = reader.len();
314 let tag: u32 = <u32 as #felt_repr_crate::FromFeltRepr>::from_felt_repr(reader)?;
315 match tag {
316 #(#arms,)*
317 other => Err(#felt_repr_crate::FeltReprError::UnknownEnumTag {
318 pos: tag_pos,
319 len,
320 ty: stringify!(#name),
321 tag: other,
322 }),
323 }
324 }
325 }
326 }
327 }
328 Data::Union(_) => {
329 return Err(Error::new(
330 input.span(),
331 format!("{trait_name} cannot be derived for union `{name}`"),
332 ));
333 }
334 };
335
336 let expanded = quote! {
337 #expanded
338
339 impl #impl_generics ::core::convert::TryFrom<&[#felt_ty]> for #name #ty_generics #where_clause {
340 type Error = #felt_repr_crate::FeltReprError;
341
342 #[inline(always)]
343 fn try_from(felts: &[#felt_ty]) -> Result<Self, Self::Error> {
344 let mut reader = #felt_repr_crate::FeltReader::new(felts);
345 let value = <Self as #felt_repr_crate::FromFeltRepr>::from_felt_repr(&mut reader)?;
346 reader.ensure_eof()?;
347 Ok(value)
348 }
349 }
350 };
351
352 Ok(expanded.into())
353}
354
355#[proc_macro_derive(DeriveToFeltRepr)]
374pub fn derive_to_felt_repr(input: TokenStream) -> TokenStream {
375 let input = parse_macro_input!(input as DeriveInput);
376
377 match derive_to_felt_repr_impl(&input, quote!(miden_field_repr)) {
378 Ok(ts) => ts,
379 Err(err) => err.into_compile_error().into(),
380 }
381}
382
383fn derive_to_felt_repr_impl(
384 input: &DeriveInput,
385 felt_repr_crate: TokenStream2,
386) -> Result<TokenStream, Error> {
387 let name = &input.ident;
388 let generics = &input.generics;
389 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
390
391 let trait_name = "ToFeltRepr";
392 let expanded = match &input.data {
393 Data::Struct(_) => match extract_struct_fields(input, trait_name)? {
394 StructFields::Named(fields) => {
395 let field_names: Vec<_> =
396 fields.iter().map(|field| field.ident.as_ref().unwrap()).collect();
397 quote! {
398 impl #impl_generics #felt_repr_crate::ToFeltRepr for #name #ty_generics #where_clause {
399 fn write_felt_repr(&self, writer: &mut #felt_repr_crate::FeltWriter<'_>) {
400 #(#felt_repr_crate::ToFeltRepr::write_felt_repr(&self.#field_names, writer);)*
401 }
402 }
403 }
404 }
405 StructFields::Unnamed(fields) => {
406 let field_indexes: Vec<Index> = (0..fields.len()).map(Index::from).collect();
407 quote! {
408 impl #impl_generics #felt_repr_crate::ToFeltRepr for #name #ty_generics #where_clause {
409 fn write_felt_repr(&self, writer: &mut #felt_repr_crate::FeltWriter<'_>) {
410 #(#felt_repr_crate::ToFeltRepr::write_felt_repr(&self.#field_indexes, writer);)*
411 }
412 }
413 }
414 }
415 },
416 Data::Enum(_) => {
417 let variants = extract_enum_variants(input, trait_name)?;
418 ensure_no_explicit_discriminants(variants, trait_name, name)?;
419
420 let arms = variants.iter().enumerate().map(|(variant_ordinal, variant)| {
421 let variant_ident = &variant.ident;
422 let tag = variant_ordinal as u32;
423
424 match &variant.fields {
425 Fields::Unit => quote! {
426 Self::#variant_ident => {
427 #felt_repr_crate::ToFeltRepr::write_felt_repr(&(#tag as u32), writer);
428 return;
429 }
430 },
431 Fields::Unnamed(fields) => {
432 let bindings: Vec<_> = (0..fields.unnamed.len())
433 .map(|i| format_ident!("__field{i}"))
434 .collect();
435 quote! {
436 Self::#variant_ident(#(#bindings),*) => {
437 #felt_repr_crate::ToFeltRepr::write_felt_repr(&(#tag as u32), writer);
438 #(#felt_repr_crate::ToFeltRepr::write_felt_repr(#bindings, writer);)*
439 return;
440 }
441 }
442 }
443 Fields::Named(fields) => {
444 let bindings: Vec<_> = fields
445 .named
446 .iter()
447 .map(|f| f.ident.as_ref().expect("named field"))
448 .collect();
449 quote! {
450 Self::#variant_ident { #(#bindings),* } => {
451 #felt_repr_crate::ToFeltRepr::write_felt_repr(&(#tag as u32), writer);
452 #(#felt_repr_crate::ToFeltRepr::write_felt_repr(#bindings, writer);)*
453 return;
454 }
455 }
456 }
457 }
458 });
459
460 quote! {
461 impl #impl_generics #felt_repr_crate::ToFeltRepr for #name #ty_generics #where_clause {
462 #[inline(always)]
463 fn write_felt_repr(&self, writer: &mut #felt_repr_crate::FeltWriter<'_>) {
464 match self {
465 #(#arms,)*
466 }
467 }
468 }
469 }
470 }
471 Data::Union(_) => {
472 return Err(Error::new(
473 input.span(),
474 format!("{trait_name} cannot be derived for union `{name}`"),
475 ));
476 }
477 };
478
479 Ok(expanded.into())
480}