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