1 use proc_macro2::TokenStream;
2 use quote::{format_ident, quote};
3 use std::collections::HashSet;
4 use std::fmt;
5 use syn::parse::{Parse, ParseStream};
6 use syn::punctuated::Punctuated;
7 use syn::{braced, parse_quote, Data, DeriveInput, Error, Result, Token};
8 use wasmtime_component_util::{DiscriminantSize, FlagsSize};
9 
10 mod kw {
11     syn::custom_keyword!(record);
12     syn::custom_keyword!(variant);
13     syn::custom_keyword!(flags);
14     syn::custom_keyword!(name);
15 }
16 
17 #[derive(Debug, Copy, Clone)]
18 pub enum VariantStyle {
19     Variant,
20     Enum,
21 }
22 
23 impl fmt::Display for VariantStyle {
24     fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
25         f.write_str(match self {
26             Self::Variant => "variant",
27             Self::Enum => "enum",
28         })
29     }
30 }
31 
32 #[derive(Debug, Copy, Clone)]
33 enum Style {
34     Record,
35     Variant(VariantStyle),
36 }
37 
38 fn find_style(input: &DeriveInput) -> Result<Style> {
39     let mut style = None;
40 
41     for attribute in &input.attrs {
42         if !attribute.path().is_ident("component") {
43             continue;
44         }
45         let attr_style = attribute.parse_args()?;
46 
47         if style.is_some() {
48             return Err(Error::new_spanned(
49                 attribute,
50                 "duplicate `component` attribute",
51             ));
52         }
53         style = Some(attr_style);
54     }
55 
56     style.ok_or_else(|| Error::new_spanned(input, "missing `component` attribute"))
57 }
58 
59 impl Parse for Style {
60     fn parse(input: ParseStream) -> Result<Self> {
61         let lookahead = input.lookahead1();
62         if lookahead.peek(kw::record) {
63             input.parse::<kw::record>()?;
64             Ok(Style::Record)
65         } else if lookahead.peek(kw::variant) {
66             input.parse::<kw::variant>()?;
67             Ok(Style::Variant(VariantStyle::Variant))
68         } else if lookahead.peek(Token![enum]) {
69             input.parse::<Token![enum]>()?;
70             Ok(Style::Variant(VariantStyle::Enum))
71         } else if input.peek(kw::flags) {
72             Err(input.error(
73                 "`flags` not allowed here; \
74                  use `wasmtime::component::flags!` macro to define `flags` types",
75             ))
76         } else {
77             Err(lookahead.error())
78         }
79     }
80 }
81 
82 fn find_rename(attributes: &[syn::Attribute]) -> Result<Option<syn::LitStr>> {
83     let mut name = None;
84 
85     for attribute in attributes {
86         if !attribute.path().is_ident("component") {
87             continue;
88         }
89         let name_literal = attribute.parse_args_with(|parser: ParseStream<'_>| {
90             parser.parse::<kw::name>()?;
91             parser.parse::<Token![=]>()?;
92             parser.parse::<syn::LitStr>()
93         })?;
94 
95         if name.is_some() {
96             return Err(Error::new_spanned(
97                 attribute,
98                 "duplicate field rename attribute",
99             ));
100         }
101 
102         name = Some(name_literal);
103     }
104 
105     Ok(name)
106 }
107 
108 fn add_trait_bounds(generics: &syn::Generics, bound: syn::TypeParamBound) -> syn::Generics {
109     let mut generics = generics.clone();
110     for param in &mut generics.params {
111         if let syn::GenericParam::Type(ref mut type_param) = *param {
112             type_param.bounds.push(bound.clone());
113         }
114     }
115     generics
116 }
117 
118 pub struct VariantCase<'a> {
119     attrs: &'a [syn::Attribute],
120     ident: &'a syn::Ident,
121     ty: Option<&'a syn::Type>,
122 }
123 
124 pub trait Expander {
125     fn expand_record(
126         &self,
127         name: &syn::Ident,
128         generics: &syn::Generics,
129         fields: &[&syn::Field],
130     ) -> Result<TokenStream>;
131 
132     fn expand_variant(
133         &self,
134         name: &syn::Ident,
135         generics: &syn::Generics,
136         discriminant_size: DiscriminantSize,
137         cases: &[VariantCase],
138         style: VariantStyle,
139     ) -> Result<TokenStream>;
140 }
141 
142 pub fn expand(expander: &dyn Expander, input: &DeriveInput) -> Result<TokenStream> {
143     match find_style(input)? {
144         Style::Record => expand_record(expander, input),
145         Style::Variant(style) => expand_variant(expander, input, style),
146     }
147 }
148 
149 fn expand_record(expander: &dyn Expander, input: &DeriveInput) -> Result<TokenStream> {
150     let name = &input.ident;
151 
152     let body = if let Data::Struct(body) = &input.data {
153         body
154     } else {
155         return Err(Error::new(
156             name.span(),
157             "`record` component types can only be derived for Rust `struct`s",
158         ));
159     };
160 
161     match &body.fields {
162         syn::Fields::Named(fields) => expander.expand_record(
163             &input.ident,
164             &input.generics,
165             &fields.named.iter().collect::<Vec<_>>(),
166         ),
167 
168         syn::Fields::Unnamed(_) | syn::Fields::Unit => Err(Error::new(
169             name.span(),
170             "`record` component types can only be derived for `struct`s with named fields",
171         )),
172     }
173 }
174 
175 fn expand_variant(
176     expander: &dyn Expander,
177     input: &DeriveInput,
178     style: VariantStyle,
179 ) -> Result<TokenStream> {
180     let name = &input.ident;
181 
182     let body = if let Data::Enum(body) = &input.data {
183         body
184     } else {
185         return Err(Error::new(
186             name.span(),
187             format!(
188                 "`{}` component types can only be derived for Rust `enum`s",
189                 style
190             ),
191         ));
192     };
193 
194     if body.variants.is_empty() {
195         return Err(Error::new(
196             name.span(),
197             format!("`{}` component types can only be derived for Rust `enum`s with at least one variant", style),
198         ));
199     }
200 
201     let discriminant_size = DiscriminantSize::from_count(body.variants.len()).ok_or_else(|| {
202         Error::new(
203             input.ident.span(),
204             "`enum`s with more than 2^32 variants are not supported",
205         )
206     })?;
207 
208     let cases = body
209         .variants
210         .iter()
211         .map(
212             |syn::Variant {
213                  attrs,
214                  ident,
215                  fields,
216                  ..
217              }| {
218                 Ok(VariantCase {
219                     attrs,
220                     ident,
221                     ty: match fields {
222                         syn::Fields::Unnamed(fields) if fields.unnamed.len() == 1 => {
223                             Some(&fields.unnamed[0].ty)
224                         }
225                         syn::Fields::Unit => None,
226                         _ => {
227                             return Err(Error::new(
228                                 name.span(),
229                                 format!(
230                                     "`{}` component types can only be derived for Rust `enum`s \
231                                      containing variants with {}",
232                                     style,
233                                     match style {
234                                         VariantStyle::Variant => "at most one unnamed field each",
235                                         VariantStyle::Enum => "no fields",
236                                     }
237                                 ),
238                             ))
239                         }
240                     },
241                 })
242             },
243         )
244         .collect::<Result<Vec<_>>>()?;
245 
246     expander.expand_variant(
247         &input.ident,
248         &input.generics,
249         discriminant_size,
250         &cases,
251         style,
252     )
253 }
254 
255 fn expand_record_for_component_type(
256     name: &syn::Ident,
257     generics: &syn::Generics,
258     fields: &[&syn::Field],
259     typecheck: TokenStream,
260     typecheck_argument: TokenStream,
261 ) -> Result<TokenStream> {
262     let internal = quote!(wasmtime::component::__internal);
263 
264     let mut lower_generic_params = TokenStream::new();
265     let mut lower_generic_args = TokenStream::new();
266     let mut lower_field_declarations = TokenStream::new();
267     let mut abi_list = TokenStream::new();
268     let mut unique_types = HashSet::new();
269 
270     for (index, syn::Field { ident, ty, .. }) in fields.iter().enumerate() {
271         let generic = format_ident!("T{}", index);
272 
273         lower_generic_params.extend(quote!(#generic: Copy,));
274         lower_generic_args.extend(quote!(<#ty as wasmtime::component::ComponentType>::Lower,));
275 
276         lower_field_declarations.extend(quote!(#ident: #generic,));
277 
278         abi_list.extend(quote!(
279             <#ty as wasmtime::component::ComponentType>::ABI,
280         ));
281 
282         unique_types.insert(ty);
283     }
284 
285     let generics = add_trait_bounds(generics, parse_quote!(wasmtime::component::ComponentType));
286     let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
287     let lower = format_ident!("Lower{}", name);
288 
289     // You may wonder why we make the types of all the fields of the #lower struct generic.  This is to work
290     // around the lack of [perfect derive support in
291     // rustc](https://smallcultfollowing.com/babysteps//blog/2022/04/12/implied-bounds-and-perfect-derive/#what-is-perfect-derive)
292     // as of this writing.
293     //
294     // If the struct we're deriving a `ComponentType` impl for has any generic parameters, then #lower needs
295     // generic parameters too.  And if we just copy the parameters and bounds from the impl to #lower, then the
296     // `#[derive(Clone, Copy)]` will fail unless the original generics were declared with those bounds, which
297     // we don't want to require.
298     //
299     // Alternatively, we could just pass the `Lower` associated type of each generic type as arguments to
300     // #lower, but that would require distinguishing between generic and concrete types when generating
301     // #lower_field_declarations, which would require some form of symbol resolution.  That doesn't seem worth
302     // the trouble.
303 
304     let expanded = quote! {
305         #[doc(hidden)]
306         #[derive(Clone, Copy)]
307         #[repr(C)]
308         pub struct #lower <#lower_generic_params> {
309             #lower_field_declarations
310             _align: [wasmtime::ValRaw; 0],
311         }
312 
313         unsafe impl #impl_generics wasmtime::component::ComponentType for #name #ty_generics #where_clause {
314             type Lower = #lower <#lower_generic_args>;
315 
316             const ABI: #internal::CanonicalAbiInfo =
317                 #internal::CanonicalAbiInfo::record_static(&[#abi_list]);
318 
319             #[inline]
320             fn typecheck(
321                 ty: &#internal::InterfaceType,
322                 types: &#internal::InstanceType<'_>,
323             ) -> #internal::anyhow::Result<()> {
324                 #internal::#typecheck(ty, types, &[#typecheck_argument])
325             }
326         }
327     };
328 
329     Ok(quote!(const _: () = { #expanded };))
330 }
331 
332 fn quote(size: DiscriminantSize, discriminant: usize) -> TokenStream {
333     match size {
334         DiscriminantSize::Size1 => {
335             let discriminant = u8::try_from(discriminant).unwrap();
336             quote!(#discriminant)
337         }
338         DiscriminantSize::Size2 => {
339             let discriminant = u16::try_from(discriminant).unwrap();
340             quote!(#discriminant)
341         }
342         DiscriminantSize::Size4 => {
343             let discriminant = u32::try_from(discriminant).unwrap();
344             quote!(#discriminant)
345         }
346     }
347 }
348 
349 pub struct LiftExpander;
350 
351 impl Expander for LiftExpander {
352     fn expand_record(
353         &self,
354         name: &syn::Ident,
355         generics: &syn::Generics,
356         fields: &[&syn::Field],
357     ) -> Result<TokenStream> {
358         let internal = quote!(wasmtime::component::__internal);
359 
360         let mut lifts = TokenStream::new();
361         let mut loads = TokenStream::new();
362 
363         for (i, syn::Field { ident, ty, .. }) in fields.iter().enumerate() {
364             let field_ty = quote!(ty.fields[#i].ty);
365             lifts.extend(quote!(#ident: <#ty as wasmtime::component::Lift>::lift(
366                 cx, #field_ty, &src.#ident
367             )?,));
368 
369             loads.extend(quote!(#ident: <#ty as wasmtime::component::Lift>::load(
370                 cx, #field_ty,
371                 &bytes
372                     [<#ty as wasmtime::component::ComponentType>::ABI.next_field32_size(&mut offset)..]
373                     [..<#ty as wasmtime::component::ComponentType>::SIZE32]
374             )?,));
375         }
376 
377         let generics = add_trait_bounds(generics, parse_quote!(wasmtime::component::Lift));
378         let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
379 
380         let extract_ty = quote! {
381             let ty = match ty {
382                 #internal::InterfaceType::Record(i) => &cx.types[i],
383                 _ => #internal::bad_type_info(),
384             };
385         };
386 
387         let expanded = quote! {
388             unsafe impl #impl_generics wasmtime::component::Lift for #name #ty_generics #where_clause {
389                 #[inline]
390                 fn lift(
391                     cx: &mut #internal::LiftContext<'_>,
392                     ty: #internal::InterfaceType,
393                     src: &Self::Lower,
394                 ) -> #internal::anyhow::Result<Self> {
395                     #extract_ty
396                     Ok(Self {
397                         #lifts
398                     })
399                 }
400 
401                 #[inline]
402                 fn load(
403                     cx: &mut #internal::LiftContext<'_>,
404                     ty: #internal::InterfaceType,
405                     bytes: &[u8],
406                 ) -> #internal::anyhow::Result<Self> {
407                     #extract_ty
408                     debug_assert!(
409                         (bytes.as_ptr() as usize)
410                             % (<Self as wasmtime::component::ComponentType>::ALIGN32 as usize)
411                             == 0
412                     );
413                     let mut offset = 0;
414                     Ok(Self {
415                         #loads
416                     })
417                 }
418             }
419         };
420 
421         Ok(expanded)
422     }
423 
424     fn expand_variant(
425         &self,
426         name: &syn::Ident,
427         generics: &syn::Generics,
428         discriminant_size: DiscriminantSize,
429         cases: &[VariantCase],
430         style: VariantStyle,
431     ) -> Result<TokenStream> {
432         let internal = quote!(wasmtime::component::__internal);
433 
434         let mut lifts = TokenStream::new();
435         let mut loads = TokenStream::new();
436 
437         let interface_type_variant = match style {
438             VariantStyle::Variant => quote!(Variant),
439             VariantStyle::Enum => quote!(Enum),
440         };
441 
442         for (index, VariantCase { ident, ty, .. }) in cases.iter().enumerate() {
443             let index_u32 = u32::try_from(index).unwrap();
444 
445             let index_quoted = quote(discriminant_size, index);
446 
447             if let Some(ty) = ty {
448                 let payload_ty = match style {
449                     VariantStyle::Variant => {
450                         quote!(ty.cases[#index].unwrap_or_else(#internal::bad_type_info))
451                     }
452                     VariantStyle::Enum => unreachable!(),
453                 };
454                 lifts.extend(
455                     quote!(#index_u32 => Self::#ident(<#ty as wasmtime::component::Lift>::lift(
456                         cx, #payload_ty, unsafe { &src.payload.#ident }
457                     )?),),
458                 );
459 
460                 loads.extend(
461                     quote!(#index_quoted => Self::#ident(<#ty as wasmtime::component::Lift>::load(
462                         cx, #payload_ty, &payload[..<#ty as wasmtime::component::ComponentType>::SIZE32]
463                     )?),),
464                 );
465             } else {
466                 lifts.extend(quote!(#index_u32 => Self::#ident,));
467 
468                 loads.extend(quote!(#index_quoted => Self::#ident,));
469             }
470         }
471 
472         let generics = add_trait_bounds(generics, parse_quote!(wasmtime::component::Lift));
473         let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
474 
475         let from_bytes = match discriminant_size {
476             DiscriminantSize::Size1 => quote!(bytes[0]),
477             DiscriminantSize::Size2 => quote!(u16::from_le_bytes(bytes[0..2].try_into()?)),
478             DiscriminantSize::Size4 => quote!(u32::from_le_bytes(bytes[0..4].try_into()?)),
479         };
480 
481         let extract_ty = quote! {
482             let ty = match ty {
483                 #internal::InterfaceType::#interface_type_variant(i) => &cx.types[i],
484                 _ => #internal::bad_type_info(),
485             };
486         };
487 
488         let expanded = quote! {
489             unsafe impl #impl_generics wasmtime::component::Lift for #name #ty_generics #where_clause {
490                 #[inline]
491                 fn lift(
492                     cx: &mut #internal::LiftContext<'_>,
493                     ty: #internal::InterfaceType,
494                     src: &Self::Lower,
495                 ) -> #internal::anyhow::Result<Self> {
496                     #extract_ty
497                     Ok(match src.tag.get_u32() {
498                         #lifts
499                         discrim => #internal::anyhow::bail!("unexpected discriminant: {}", discrim),
500                     })
501                 }
502 
503                 #[inline]
504                 fn load(
505                     cx: &mut #internal::LiftContext<'_>,
506                     ty: #internal::InterfaceType,
507                     bytes: &[u8],
508                 ) -> #internal::anyhow::Result<Self> {
509                     let align = <Self as wasmtime::component::ComponentType>::ALIGN32;
510                     debug_assert!((bytes.as_ptr() as usize) % (align as usize) == 0);
511                     let discrim = #from_bytes;
512                     let payload_offset = <Self as #internal::ComponentVariant>::PAYLOAD_OFFSET32;
513                     let payload = &bytes[payload_offset..];
514                     #extract_ty
515                     Ok(match discrim {
516                         #loads
517                         discrim => #internal::anyhow::bail!("unexpected discriminant: {}", discrim),
518                     })
519                 }
520             }
521         };
522 
523         Ok(expanded)
524     }
525 }
526 
527 pub struct LowerExpander;
528 
529 impl Expander for LowerExpander {
530     fn expand_record(
531         &self,
532         name: &syn::Ident,
533         generics: &syn::Generics,
534         fields: &[&syn::Field],
535     ) -> Result<TokenStream> {
536         let internal = quote!(wasmtime::component::__internal);
537 
538         let mut lowers = TokenStream::new();
539         let mut stores = TokenStream::new();
540 
541         for (i, syn::Field { ident, ty, .. }) in fields.iter().enumerate() {
542             let field_ty = quote!(ty.fields[#i].ty);
543             lowers.extend(quote!(wasmtime::component::Lower::lower(
544                 &self.#ident, cx, #field_ty, #internal::map_maybe_uninit!(dst.#ident)
545             )?;));
546 
547             stores.extend(quote!(wasmtime::component::Lower::store(
548                 &self.#ident,
549                 cx,
550                 #field_ty,
551                 <#ty as wasmtime::component::ComponentType>::ABI.next_field32_size(&mut offset),
552             )?;));
553         }
554 
555         let generics = add_trait_bounds(generics, parse_quote!(wasmtime::component::Lower));
556         let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
557 
558         let extract_ty = quote! {
559             let ty = match ty {
560                 #internal::InterfaceType::Record(i) => &cx.types[i],
561                 _ => #internal::bad_type_info(),
562             };
563         };
564 
565         let expanded = quote! {
566             unsafe impl #impl_generics wasmtime::component::Lower for #name #ty_generics #where_clause {
567                 #[inline]
568                 fn lower<T>(
569                     &self,
570                     cx: &mut #internal::LowerContext<'_, T>,
571                     ty: #internal::InterfaceType,
572                     dst: &mut core::mem::MaybeUninit<Self::Lower>,
573                 ) -> #internal::anyhow::Result<()> {
574                     #extract_ty
575                     #lowers
576                     Ok(())
577                 }
578 
579                 #[inline]
580                 fn store<T>(
581                     &self,
582                     cx: &mut #internal::LowerContext<'_, T>,
583                     ty: #internal::InterfaceType,
584                     mut offset: usize
585                 ) -> #internal::anyhow::Result<()> {
586                     debug_assert!(offset % (<Self as wasmtime::component::ComponentType>::ALIGN32 as usize) == 0);
587                     #extract_ty
588                     #stores
589                     Ok(())
590                 }
591             }
592         };
593 
594         Ok(expanded)
595     }
596 
597     fn expand_variant(
598         &self,
599         name: &syn::Ident,
600         generics: &syn::Generics,
601         discriminant_size: DiscriminantSize,
602         cases: &[VariantCase],
603         style: VariantStyle,
604     ) -> Result<TokenStream> {
605         let internal = quote!(wasmtime::component::__internal);
606 
607         let mut lowers = TokenStream::new();
608         let mut stores = TokenStream::new();
609 
610         let interface_type_variant = match style {
611             VariantStyle::Variant => quote!(Variant),
612             VariantStyle::Enum => quote!(Enum),
613         };
614 
615         for (index, VariantCase { ident, ty, .. }) in cases.iter().enumerate() {
616             let index_u32 = u32::try_from(index).unwrap();
617 
618             let index_quoted = quote(discriminant_size, index);
619 
620             let discriminant_size = usize::from(discriminant_size);
621 
622             let pattern;
623             let lower;
624             let store;
625 
626             if ty.is_some() {
627                 let ty = match style {
628                     VariantStyle::Variant => {
629                         quote!(ty.cases[#index].unwrap_or_else(#internal::bad_type_info))
630                     }
631                     VariantStyle::Enum => unreachable!(),
632                 };
633                 pattern = quote!(Self::#ident(value));
634                 lower = quote!(value.lower(cx, #ty, dst));
635                 store = quote!(value.store(
636                     cx,
637                     #ty,
638                     offset + <Self as #internal::ComponentVariant>::PAYLOAD_OFFSET32,
639                 ));
640             } else {
641                 pattern = quote!(Self::#ident);
642                 lower = quote!(Ok(()));
643                 store = quote!(Ok(()));
644             }
645 
646             lowers.extend(quote!(#pattern => {
647                 #internal::map_maybe_uninit!(dst.tag).write(wasmtime::ValRaw::u32(#index_u32));
648                 unsafe {
649                     #internal::lower_payload(
650                         #internal::map_maybe_uninit!(dst.payload),
651                         |payload| #internal::map_maybe_uninit!(payload.#ident),
652                         |dst| #lower,
653                     )
654                 }
655             }));
656 
657             stores.extend(quote!(#pattern => {
658                 *cx.get::<#discriminant_size>(offset) = #index_quoted.to_le_bytes();
659                 #store
660             }));
661         }
662 
663         let generics = add_trait_bounds(generics, parse_quote!(wasmtime::component::Lower));
664         let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
665 
666         let extract_ty = quote! {
667             let ty = match ty {
668                 #internal::InterfaceType::#interface_type_variant(i) => &cx.types[i],
669                 _ => #internal::bad_type_info(),
670             };
671         };
672 
673         let expanded = quote! {
674             unsafe impl #impl_generics wasmtime::component::Lower for #name #ty_generics #where_clause {
675                 #[inline]
676                 fn lower<T>(
677                     &self,
678                     cx: &mut #internal::LowerContext<'_, T>,
679                     ty: #internal::InterfaceType,
680                     dst: &mut core::mem::MaybeUninit<Self::Lower>,
681                 ) -> #internal::anyhow::Result<()> {
682                     #extract_ty
683                     match self {
684                         #lowers
685                     }
686                 }
687 
688                 #[inline]
689                 fn store<T>(
690                     &self,
691                     cx: &mut #internal::LowerContext<'_, T>,
692                     ty: #internal::InterfaceType,
693                     mut offset: usize
694                 ) -> #internal::anyhow::Result<()> {
695                     #extract_ty
696                     debug_assert!(offset % (<Self as wasmtime::component::ComponentType>::ALIGN32 as usize) == 0);
697                     match self {
698                         #stores
699                     }
700                 }
701             }
702         };
703 
704         Ok(expanded)
705     }
706 }
707 
708 pub struct ComponentTypeExpander;
709 
710 impl Expander for ComponentTypeExpander {
711     fn expand_record(
712         &self,
713         name: &syn::Ident,
714         generics: &syn::Generics,
715         fields: &[&syn::Field],
716     ) -> Result<TokenStream> {
717         expand_record_for_component_type(
718             name,
719             generics,
720             fields,
721             quote!(typecheck_record),
722             fields
723                 .iter()
724                 .map(
725                     |syn::Field {
726                          attrs, ident, ty, ..
727                      }| {
728                         let name = find_rename(attrs)?.unwrap_or_else(|| {
729                             let ident = ident.as_ref().unwrap();
730                             syn::LitStr::new(&ident.to_string(), ident.span())
731                         });
732 
733                         Ok(quote!((#name, <#ty as wasmtime::component::ComponentType>::typecheck),))
734                     },
735                 )
736                 .collect::<Result<_>>()?,
737         )
738     }
739 
740     fn expand_variant(
741         &self,
742         name: &syn::Ident,
743         generics: &syn::Generics,
744         _discriminant_size: DiscriminantSize,
745         cases: &[VariantCase],
746         style: VariantStyle,
747     ) -> Result<TokenStream> {
748         let internal = quote!(wasmtime::component::__internal);
749 
750         let mut case_names_and_checks = TokenStream::new();
751         let mut lower_payload_generic_params = TokenStream::new();
752         let mut lower_payload_generic_args = TokenStream::new();
753         let mut lower_payload_case_declarations = TokenStream::new();
754         let mut lower_generic_args = TokenStream::new();
755         let mut abi_list = TokenStream::new();
756         let mut unique_types = HashSet::new();
757 
758         for (index, VariantCase { attrs, ident, ty }) in cases.iter().enumerate() {
759             let rename = find_rename(attrs)?;
760 
761             let name = rename.unwrap_or_else(|| syn::LitStr::new(&ident.to_string(), ident.span()));
762 
763             if let Some(ty) = ty {
764                 abi_list.extend(quote!(Some(<#ty as wasmtime::component::ComponentType>::ABI),));
765 
766                 case_names_and_checks.extend(match style {
767                     VariantStyle::Variant => {
768                         quote!((#name, Some(<#ty as wasmtime::component::ComponentType>::typecheck)),)
769                     }
770                     VariantStyle::Enum => {
771                         return Err(Error::new(
772                             ident.span(),
773                             "payloads are not permitted for `enum` cases",
774                         ))
775                     }
776                 });
777 
778                 let generic = format_ident!("T{}", index);
779 
780                 lower_payload_generic_params.extend(quote!(#generic: Copy,));
781                 lower_payload_generic_args.extend(quote!(#generic,));
782                 lower_payload_case_declarations.extend(quote!(#ident: #generic,));
783                 lower_generic_args
784                     .extend(quote!(<#ty as wasmtime::component::ComponentType>::Lower,));
785 
786                 unique_types.insert(ty);
787             } else {
788                 abi_list.extend(quote!(None,));
789                 case_names_and_checks.extend(match style {
790                     VariantStyle::Variant => {
791                         quote!((#name, None),)
792                     }
793                     VariantStyle::Enum => quote!(#name,),
794                 });
795                 lower_payload_case_declarations.extend(quote!(#ident: [wasmtime::ValRaw; 0],));
796             }
797         }
798 
799         let typecheck = match style {
800             VariantStyle::Variant => quote!(typecheck_variant),
801             VariantStyle::Enum => quote!(typecheck_enum),
802         };
803 
804         let generics = add_trait_bounds(generics, parse_quote!(wasmtime::component::ComponentType));
805         let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
806         let lower = format_ident!("Lower{}", name);
807         let lower_payload = format_ident!("LowerPayload{}", name);
808 
809         // You may wonder why we make the types of all the fields of the #lower struct and #lower_payload union
810         // generic.  This is to work around a [normalization bug in
811         // rustc](https://github.com/rust-lang/rust/issues/90903) such that the compiler does not understand that
812         // e.g. `<i32 as ComponentType>::Lower` is `Copy` despite the bound specified in `ComponentType`'s
813         // definition.
814         //
815         // See also the comment in `Self::expand_record` above for another reason why we do this.
816 
817         let expanded = quote! {
818             #[doc(hidden)]
819             #[derive(Clone, Copy)]
820             #[repr(C)]
821             pub struct #lower<#lower_payload_generic_params> {
822                 tag: wasmtime::ValRaw,
823                 payload: #lower_payload<#lower_payload_generic_args>
824             }
825 
826             #[doc(hidden)]
827             #[allow(non_snake_case)]
828             #[derive(Clone, Copy)]
829             #[repr(C)]
830             union #lower_payload<#lower_payload_generic_params> {
831                 #lower_payload_case_declarations
832             }
833 
834             unsafe impl #impl_generics wasmtime::component::ComponentType for #name #ty_generics #where_clause {
835                 type Lower = #lower<#lower_generic_args>;
836 
837                 #[inline]
838                 fn typecheck(
839                     ty: &#internal::InterfaceType,
840                     types: &#internal::InstanceType<'_>,
841                 ) -> #internal::anyhow::Result<()> {
842                     #internal::#typecheck(ty, types, &[#case_names_and_checks])
843                 }
844 
845                 const ABI: #internal::CanonicalAbiInfo =
846                     #internal::CanonicalAbiInfo::variant_static(&[#abi_list]);
847             }
848 
849             unsafe impl #impl_generics #internal::ComponentVariant for #name #ty_generics #where_clause {
850                 const CASES: &'static [Option<#internal::CanonicalAbiInfo>] = &[#abi_list];
851             }
852         };
853 
854         Ok(quote!(const _: () = { #expanded };))
855     }
856 }
857 
858 #[derive(Debug)]
859 struct Flag {
860     rename: Option<String>,
861     name: String,
862 }
863 
864 impl Parse for Flag {
865     fn parse(input: ParseStream) -> Result<Self> {
866         let attributes = syn::Attribute::parse_outer(input)?;
867 
868         let rename = find_rename(&attributes)?.map(|literal| literal.value());
869 
870         input.parse::<Token![const]>()?;
871         let name = input.parse::<syn::Ident>()?.to_string();
872 
873         Ok(Self { rename, name })
874     }
875 }
876 
877 #[derive(Debug)]
878 pub struct Flags {
879     name: String,
880     flags: Vec<Flag>,
881 }
882 
883 impl Parse for Flags {
884     fn parse(input: ParseStream) -> Result<Self> {
885         let name = input.parse::<syn::Ident>()?.to_string();
886 
887         let content;
888         braced!(content in input);
889 
890         let flags = content
891             .parse_terminated(Flag::parse, Token![;])?
892             .into_iter()
893             .collect();
894 
895         Ok(Self { name, flags })
896     }
897 }
898 
899 pub fn expand_flags(flags: &Flags) -> Result<TokenStream> {
900     let size = FlagsSize::from_count(flags.flags.len());
901 
902     let ty;
903     let eq;
904 
905     let count = flags.flags.len();
906 
907     match size {
908         FlagsSize::Size0 => {
909             ty = quote!(());
910             eq = quote!(true);
911         }
912         FlagsSize::Size1 => {
913             ty = quote!(u8);
914 
915             eq = if count == 8 {
916                 quote!(self.__inner0.eq(&rhs.__inner0))
917             } else {
918                 let mask = !(0xFF_u8 << count);
919 
920                 quote!((self.__inner0 & #mask).eq(&(rhs.__inner0 & #mask)))
921             };
922         }
923         FlagsSize::Size2 => {
924             ty = quote!(u16);
925 
926             eq = if count == 16 {
927                 quote!(self.__inner0.eq(&rhs.__inner0))
928             } else {
929                 let mask = !(0xFFFF_u16 << count);
930 
931                 quote!((self.__inner0 & #mask).eq(&(rhs.__inner0 & #mask)))
932             };
933         }
934         FlagsSize::Size4Plus(n) => {
935             ty = quote!(u32);
936 
937             let comparisons = (0..(n - 1))
938                 .map(|index| {
939                     let field = format_ident!("__inner{}", index);
940 
941                     quote!(self.#field.eq(&rhs.#field) &&)
942                 })
943                 .collect::<TokenStream>();
944 
945             let field = format_ident!("__inner{}", n - 1);
946 
947             eq = if count % 32 == 0 {
948                 quote!(#comparisons self.#field.eq(&rhs.#field))
949             } else {
950                 let mask = !(0xFFFF_FFFF_u32 << (count % 32));
951 
952                 quote!(#comparisons (self.#field & #mask).eq(&(rhs.#field & #mask)))
953             }
954         }
955     }
956 
957     let count;
958     let mut as_array;
959     let mut bitor;
960     let mut bitor_assign;
961     let mut bitand;
962     let mut bitand_assign;
963     let mut bitxor;
964     let mut bitxor_assign;
965     let mut not;
966 
967     match size {
968         FlagsSize::Size0 => {
969             count = 0;
970             as_array = quote!([]);
971             bitor = quote!(Self {});
972             bitor_assign = quote!();
973             bitand = quote!(Self {});
974             bitand_assign = quote!();
975             bitxor = quote!(Self {});
976             bitxor_assign = quote!();
977             not = quote!(Self {});
978         }
979         FlagsSize::Size1 | FlagsSize::Size2 => {
980             count = 1;
981             as_array = quote!([self.__inner0 as u32]);
982             bitor = quote!(Self {
983                 __inner0: self.__inner0.bitor(rhs.__inner0)
984             });
985             bitor_assign = quote!(self.__inner0.bitor_assign(rhs.__inner0));
986             bitand = quote!(Self {
987                 __inner0: self.__inner0.bitand(rhs.__inner0)
988             });
989             bitand_assign = quote!(self.__inner0.bitand_assign(rhs.__inner0));
990             bitxor = quote!(Self {
991                 __inner0: self.__inner0.bitxor(rhs.__inner0)
992             });
993             bitxor_assign = quote!(self.__inner0.bitxor_assign(rhs.__inner0));
994             not = quote!(Self {
995                 __inner0: self.__inner0.not()
996             });
997         }
998         FlagsSize::Size4Plus(n) => {
999             count = usize::from(n);
1000             as_array = TokenStream::new();
1001             bitor = TokenStream::new();
1002             bitor_assign = TokenStream::new();
1003             bitand = TokenStream::new();
1004             bitand_assign = TokenStream::new();
1005             bitxor = TokenStream::new();
1006             bitxor_assign = TokenStream::new();
1007             not = TokenStream::new();
1008 
1009             for index in 0..n {
1010                 let field = format_ident!("__inner{}", index);
1011 
1012                 as_array.extend(quote!(self.#field,));
1013                 bitor.extend(quote!(#field: self.#field.bitor(rhs.#field),));
1014                 bitor_assign.extend(quote!(self.#field.bitor_assign(rhs.#field);));
1015                 bitand.extend(quote!(#field: self.#field.bitand(rhs.#field),));
1016                 bitand_assign.extend(quote!(self.#field.bitand_assign(rhs.#field);));
1017                 bitxor.extend(quote!(#field: self.#field.bitxor(rhs.#field),));
1018                 bitxor_assign.extend(quote!(self.#field.bitxor_assign(rhs.#field);));
1019                 not.extend(quote!(#field: self.#field.not(),));
1020             }
1021 
1022             as_array = quote!([#as_array]);
1023             bitor = quote!(Self { #bitor });
1024             bitand = quote!(Self { #bitand });
1025             bitxor = quote!(Self { #bitxor });
1026             not = quote!(Self { #not });
1027         }
1028     };
1029 
1030     let name = format_ident!("{}", flags.name);
1031 
1032     let mut constants = TokenStream::new();
1033     let mut rust_names = TokenStream::new();
1034     let mut component_names = TokenStream::new();
1035 
1036     for (index, Flag { name, rename }) in flags.flags.iter().enumerate() {
1037         rust_names.extend(quote!(#name,));
1038 
1039         let component_name = rename.as_ref().unwrap_or(name);
1040         component_names.extend(quote!(#component_name,));
1041 
1042         let fields = match size {
1043             FlagsSize::Size0 => quote!(),
1044             FlagsSize::Size1 => {
1045                 let init = 1_u8 << index;
1046                 quote!(__inner0: #init)
1047             }
1048             FlagsSize::Size2 => {
1049                 let init = 1_u16 << index;
1050                 quote!(__inner0: #init)
1051             }
1052             FlagsSize::Size4Plus(n) => (0..n)
1053                 .map(|i| {
1054                     let field = format_ident!("__inner{}", i);
1055 
1056                     let init = if index / 32 == usize::from(i) {
1057                         1_u32 << (index % 32)
1058                     } else {
1059                         0
1060                     };
1061 
1062                     quote!(#field: #init,)
1063                 })
1064                 .collect::<TokenStream>(),
1065         };
1066 
1067         let name = format_ident!("{}", name);
1068 
1069         constants.extend(quote!(pub const #name: Self = Self { #fields };));
1070     }
1071 
1072     let generics = syn::Generics {
1073         lt_token: None,
1074         params: Punctuated::new(),
1075         gt_token: None,
1076         where_clause: None,
1077     };
1078 
1079     let fields = {
1080         let ty = syn::parse2::<syn::Type>(ty.clone())?;
1081 
1082         (0..count)
1083             .map(|index| syn::Field {
1084                 attrs: Vec::new(),
1085                 vis: syn::Visibility::Inherited,
1086                 ident: Some(format_ident!("__inner{}", index)),
1087                 colon_token: None,
1088                 ty: ty.clone(),
1089                 mutability: syn::FieldMutability::None,
1090             })
1091             .collect::<Vec<_>>()
1092     };
1093 
1094     let fields = fields.iter().collect::<Vec<_>>();
1095 
1096     let component_type_impl = expand_record_for_component_type(
1097         &name,
1098         &generics,
1099         &fields,
1100         quote!(typecheck_flags),
1101         component_names,
1102     )?;
1103 
1104     let internal = quote!(wasmtime::component::__internal);
1105 
1106     let field_names = fields
1107         .iter()
1108         .map(|syn::Field { ident, .. }| ident)
1109         .collect::<Vec<_>>();
1110 
1111     let fields = fields
1112         .iter()
1113         .map(|syn::Field { ident, .. }| quote!(#[doc(hidden)] #ident: #ty,))
1114         .collect::<TokenStream>();
1115 
1116     let (field_interface_type, field_size) = match size {
1117         FlagsSize::Size0 => (quote!(NOT USED), 0usize),
1118         FlagsSize::Size1 => (quote!(#internal::InterfaceType::U8), 1),
1119         FlagsSize::Size2 => (quote!(#internal::InterfaceType::U16), 2),
1120         FlagsSize::Size4Plus(_) => (quote!(#internal::InterfaceType::U32), 4),
1121     };
1122 
1123     let expanded = quote! {
1124         #[derive(Copy, Clone, Default)]
1125         pub struct #name { #fields }
1126 
1127         impl #name {
1128             #constants
1129 
1130             pub fn as_array(&self) -> [u32; #count] {
1131                 #as_array
1132             }
1133 
1134             pub fn empty() -> Self {
1135                 Self::default()
1136             }
1137 
1138             pub fn all() -> Self {
1139                 use core::ops::Not;
1140                 Self::default().not()
1141             }
1142 
1143             pub fn contains(&self, other: Self) -> bool {
1144                 *self & other == other
1145             }
1146 
1147             pub fn intersects(&self, other: Self) -> bool {
1148                 *self & other != Self::empty()
1149             }
1150         }
1151 
1152         impl core::cmp::PartialEq for #name {
1153             fn eq(&self, rhs: &#name) -> bool {
1154                 #eq
1155             }
1156         }
1157 
1158         impl core::cmp::Eq for #name { }
1159 
1160         impl core::fmt::Debug for #name {
1161             fn fmt(&self, f: &mut core::fmt::Formatter) -> core::fmt::Result {
1162                 #internal::format_flags(&self.as_array(), &[#rust_names], f)
1163             }
1164         }
1165 
1166         impl core::ops::BitOr for #name {
1167             type Output = #name;
1168 
1169             fn bitor(self, rhs: #name) -> #name {
1170                 #bitor
1171             }
1172         }
1173 
1174         impl core::ops::BitOrAssign for #name {
1175             fn bitor_assign(&mut self, rhs: #name) {
1176                 #bitor_assign
1177             }
1178         }
1179 
1180         impl core::ops::BitAnd for #name {
1181             type Output = #name;
1182 
1183             fn bitand(self, rhs: #name) -> #name {
1184                 #bitand
1185             }
1186         }
1187 
1188         impl core::ops::BitAndAssign for #name {
1189             fn bitand_assign(&mut self, rhs: #name) {
1190                 #bitand_assign
1191             }
1192         }
1193 
1194         impl core::ops::BitXor for #name {
1195             type Output = #name;
1196 
1197             fn bitxor(self, rhs: #name) -> #name {
1198                 #bitxor
1199             }
1200         }
1201 
1202         impl core::ops::BitXorAssign for #name {
1203             fn bitxor_assign(&mut self, rhs: #name) {
1204                 #bitxor_assign
1205             }
1206         }
1207 
1208         impl core::ops::Not for #name {
1209             type Output = #name;
1210 
1211             fn not(self) -> #name {
1212                 #not
1213             }
1214         }
1215 
1216         #component_type_impl
1217 
1218         unsafe impl wasmtime::component::Lower for #name {
1219             fn lower<T>(
1220                 &self,
1221                 cx: &mut #internal::LowerContext<'_, T>,
1222                 _ty: #internal::InterfaceType,
1223                 dst: &mut core::mem::MaybeUninit<Self::Lower>,
1224             ) -> #internal::anyhow::Result<()> {
1225                 #(
1226                     self.#field_names.lower(
1227                         cx,
1228                         #field_interface_type,
1229                         #internal::map_maybe_uninit!(dst.#field_names),
1230                     )?;
1231                 )*
1232                 Ok(())
1233             }
1234 
1235             fn store<T>(
1236                 &self,
1237                 cx: &mut #internal::LowerContext<'_, T>,
1238                 _ty: #internal::InterfaceType,
1239                 mut offset: usize
1240             ) -> #internal::anyhow::Result<()> {
1241                 debug_assert!(offset % (<Self as wasmtime::component::ComponentType>::ALIGN32 as usize) == 0);
1242                 #(
1243                     self.#field_names.store(
1244                         cx,
1245                         #field_interface_type,
1246                         offset,
1247                     )?;
1248                     offset += core::mem::size_of_val(&self.#field_names);
1249                 )*
1250                 Ok(())
1251             }
1252         }
1253 
1254         unsafe impl wasmtime::component::Lift for #name {
1255             fn lift(
1256                 cx: &mut #internal::LiftContext<'_>,
1257                 _ty: #internal::InterfaceType,
1258                 src: &Self::Lower,
1259             ) -> #internal::anyhow::Result<Self> {
1260                 Ok(Self {
1261                     #(
1262                         #field_names: wasmtime::component::Lift::lift(
1263                             cx,
1264                             #field_interface_type,
1265                             &src.#field_names,
1266                         )?,
1267                     )*
1268                 })
1269             }
1270 
1271             fn load(
1272                 cx: &mut #internal::LiftContext<'_>,
1273                 _ty: #internal::InterfaceType,
1274                 bytes: &[u8],
1275             ) -> #internal::anyhow::Result<Self> {
1276                 debug_assert!(
1277                     (bytes.as_ptr() as usize)
1278                         % (<Self as wasmtime::component::ComponentType>::ALIGN32 as usize)
1279                         == 0
1280                 );
1281                 #(
1282                     let (field, bytes) = bytes.split_at(#field_size);
1283                     let #field_names = wasmtime::component::Lift::load(
1284                         cx,
1285                         #field_interface_type,
1286                         field,
1287                     )?;
1288                 )*
1289                 Ok(Self { #(#field_names,)* })
1290             }
1291         }
1292     };
1293 
1294     Ok(expanded)
1295 }
1296