1 use proc_macro2::{Span, TokenStream};
2 use quote::ToTokens;
3 use std::collections::HashMap;
4 use std::collections::HashSet;
5 use std::env;
6 use std::path::{Path, PathBuf};
7 use std::sync::atomic::{AtomicUsize, Ordering::Relaxed};
8 use syn::parse::{Error, Parse, ParseStream, Result};
9 use syn::punctuated::Punctuated;
10 use syn::{braced, token, Token};
11 use wasmtime_wit_bindgen::{AsyncConfig, Opts, Ownership, TrappableError, TrappableImports};
12 use wit_parser::{PackageId, Resolve, UnresolvedPackageGroup, WorldId};
13 
14 pub struct Config {
15     opts: Opts,
16     resolve: Resolve,
17     world: WorldId,
18     files: Vec<PathBuf>,
19     include_generated_code_from_file: bool,
20 }
21 
22 pub fn expand(input: &Config) -> Result<TokenStream> {
23     if !cfg!(feature = "async") && input.opts.async_.maybe_async() {
24         return Err(Error::new(
25             Span::call_site(),
26             "cannot enable async bindings unless `async` crate feature is active",
27         ));
28     }
29 
30     let mut src = match input.opts.generate(&input.resolve, input.world) {
31         Ok(s) => s,
32         Err(e) => return Err(Error::new(Span::call_site(), e.to_string())),
33     };
34 
35     if input.opts.stringify {
36         return Ok(quote::quote!(#src));
37     }
38 
39     // If a magical `WASMTIME_DEBUG_BINDGEN` environment variable is set then
40     // place a formatted version of the expanded code into a file. This file
41     // will then show up in rustc error messages for any codegen issues and can
42     // be inspected manually.
43     if input.include_generated_code_from_file || std::env::var("WASMTIME_DEBUG_BINDGEN").is_ok() {
44         static INVOCATION: AtomicUsize = AtomicUsize::new(0);
45         let root = Path::new(env!("DEBUG_OUTPUT_DIR"));
46         let world_name = &input.resolve.worlds[input.world].name;
47         let n = INVOCATION.fetch_add(1, Relaxed);
48         let path = root.join(format!("{world_name}{n}.rs"));
49 
50         std::fs::write(&path, &src).unwrap();
51 
52         // optimistically format the code but don't require success
53         drop(
54             std::process::Command::new("rustfmt")
55                 .arg(&path)
56                 .arg("--edition=2021")
57                 .output(),
58         );
59 
60         src = format!("include!({path:?});");
61     }
62     let mut contents = src.parse::<TokenStream>().unwrap();
63 
64     // Include a dummy `include_str!` for any files we read so rustc knows that
65     // we depend on the contents of those files.
66     for file in input.files.iter() {
67         contents.extend(
68             format!("const _: &str = include_str!(r#\"{}\"#);\n", file.display())
69                 .parse::<TokenStream>()
70                 .unwrap(),
71         );
72     }
73 
74     Ok(contents)
75 }
76 
77 impl Parse for Config {
78     fn parse(input: ParseStream<'_>) -> Result<Self> {
79         let call_site = Span::call_site();
80         let mut opts = Opts::default();
81         let mut world = None;
82         let mut inline = None;
83         let mut paths = Vec::new();
84         let mut async_configured = false;
85         let mut features = Vec::new();
86         let mut include_generated_code_from_file = false;
87 
88         if input.peek(token::Brace) {
89             let content;
90             syn::braced!(content in input);
91             let fields = Punctuated::<Opt, Token![,]>::parse_terminated(&content)?;
92             for field in fields.into_pairs() {
93                 match field.into_value() {
94                     Opt::Path(p) => {
95                         paths.extend(p.into_iter().map(|p| p.value()));
96                     }
97                     Opt::World(s) => {
98                         if world.is_some() {
99                             return Err(Error::new(s.span(), "cannot specify second world"));
100                         }
101                         world = Some(s.value());
102                     }
103                     Opt::Inline(s) => {
104                         if inline.is_some() {
105                             return Err(Error::new(s.span(), "cannot specify second source"));
106                         }
107                         inline = Some(s.value());
108                     }
109                     Opt::Tracing(val) => opts.tracing = val,
110                     Opt::VerboseTracing(val) => opts.verbose_tracing = val,
111                     Opt::Async(val, span) => {
112                         if async_configured {
113                             return Err(Error::new(span, "cannot specify second async config"));
114                         }
115                         async_configured = true;
116                         opts.async_ = val;
117                     }
118                     Opt::TrappableErrorType(val) => opts.trappable_error_type = val,
119                     Opt::TrappableImports(val) => opts.trappable_imports = val,
120                     Opt::Ownership(val) => opts.ownership = val,
121                     Opt::Interfaces(s) => {
122                         if inline.is_some() {
123                             return Err(Error::new(s.span(), "cannot specify a second source"));
124                         }
125                         inline = Some(format!(
126                             "
127                                 package wasmtime:component-macro-synthesized;
128 
129                                 world interfaces {{
130                                     {}
131                                 }}
132                             ",
133                             s.value()
134                         ));
135 
136                         if world.is_some() {
137                             return Err(Error::new(
138                                 s.span(),
139                                 "cannot specify a world with `interfaces`",
140                             ));
141                         }
142                         world = Some("interfaces".to_string());
143 
144                         opts.only_interfaces = true;
145                     }
146                     Opt::With(val) => opts.with.extend(val),
147                     Opt::AdditionalDerives(paths) => {
148                         opts.additional_derive_attributes = paths
149                             .into_iter()
150                             .map(|p| p.into_token_stream().to_string())
151                             .collect()
152                     }
153                     Opt::Stringify(val) => opts.stringify = val,
154                     Opt::SkipMutForwardingImpls(val) => opts.skip_mut_forwarding_impls = val,
155                     Opt::Features(f) => {
156                         features.extend(f.into_iter().map(|f| f.value()));
157                     }
158                     Opt::RequireStoreDataSend(val) => opts.require_store_data_send = val,
159                     Opt::WasmtimeCrate(f) => {
160                         opts.wasmtime_crate = Some(f.into_token_stream().to_string())
161                     }
162                     Opt::IncludeGeneratedCodeFromFile(i) => include_generated_code_from_file = i,
163                 }
164             }
165         } else {
166             world = input.parse::<Option<syn::LitStr>>()?.map(|s| s.value());
167             if input.parse::<Option<syn::token::In>>()?.is_some() {
168                 paths.push(input.parse::<syn::LitStr>()?.value());
169             }
170         }
171         let (resolve, pkgs, files) = parse_source(&paths, &inline, &features)
172             .map_err(|err| Error::new(call_site, format!("{err:?}")))?;
173 
174         let world = select_world(&resolve, &pkgs, world.as_deref())
175             .map_err(|e| Error::new(call_site, format!("{e:?}")))?;
176         Ok(Config {
177             opts,
178             resolve,
179             world,
180             files,
181             include_generated_code_from_file,
182         })
183     }
184 }
185 
186 fn parse_source(
187     paths: &Vec<String>,
188     inline: &Option<String>,
189     features: &[String],
190 ) -> anyhow::Result<(Resolve, Vec<PackageId>, Vec<PathBuf>)> {
191     let mut resolve = Resolve::default();
192     resolve.features.extend(features.iter().cloned());
193     let mut files = Vec::new();
194     let mut pkgs = Vec::new();
195     let root = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap());
196 
197     let parse = |resolve: &mut Resolve,
198                  files: &mut Vec<PathBuf>,
199                  pkgs: &mut Vec<PackageId>,
200                  paths: &[String]|
201      -> anyhow::Result<_> {
202         for path in paths {
203             let p = root.join(path);
204             // Try to normalize the path to make the error message more understandable when
205             // the path is not correct. Fallback to the original path if normalization fails
206             // (probably return an error somewhere else).
207             let normalized_path = match std::fs::canonicalize(&p) {
208                 Ok(p) => p,
209                 Err(_) => p.to_path_buf(),
210             };
211             let (pkg, sources) = resolve.push_path(normalized_path)?;
212             pkgs.push(pkg);
213             files.extend(sources);
214         }
215         Ok(())
216     };
217 
218     if !paths.is_empty() {
219         parse(&mut resolve, &mut files, &mut pkgs, &paths)?;
220     }
221 
222     if let Some(inline) = inline {
223         pkgs.push(resolve.push_group(UnresolvedPackageGroup::parse("macro-input", inline)?)?);
224     }
225 
226     if pkgs.is_empty() {
227         parse(&mut resolve, &mut files, &mut pkgs, &["wit".into()])?;
228     }
229 
230     Ok((resolve, pkgs, files))
231 }
232 
233 fn select_world(
234     resolve: &Resolve,
235     pkgs: &[PackageId],
236     world: Option<&str>,
237 ) -> anyhow::Result<WorldId> {
238     if pkgs.len() == 1 {
239         resolve.select_world(pkgs[0], world)
240     } else {
241         assert!(!pkgs.is_empty());
242         match world {
243             Some(name) => {
244                 if !name.contains(":") {
245                     anyhow::bail!(
246                         "with multiple packages a fully qualified \
247                          world name must be specified"
248                     )
249                 }
250 
251                 // This will ignore the package argument due to the fully
252                 // qualified name being used.
253                 resolve.select_world(pkgs[0], world)
254             }
255             None => {
256                 let worlds = pkgs
257                     .iter()
258                     .filter_map(|p| resolve.select_world(*p, None).ok())
259                     .collect::<Vec<_>>();
260                 match &worlds[..] {
261                     [] => anyhow::bail!("no packages have a world"),
262                     [world] => Ok(*world),
263                     _ => anyhow::bail!("multiple packages have a world, must specify which to use"),
264                 }
265             }
266         }
267     }
268 }
269 
270 mod kw {
271     syn::custom_keyword!(inline);
272     syn::custom_keyword!(path);
273     syn::custom_keyword!(tracing);
274     syn::custom_keyword!(verbose_tracing);
275     syn::custom_keyword!(trappable_error_type);
276     syn::custom_keyword!(world);
277     syn::custom_keyword!(ownership);
278     syn::custom_keyword!(interfaces);
279     syn::custom_keyword!(with);
280     syn::custom_keyword!(except_imports);
281     syn::custom_keyword!(only_imports);
282     syn::custom_keyword!(trappable_imports);
283     syn::custom_keyword!(additional_derives);
284     syn::custom_keyword!(stringify);
285     syn::custom_keyword!(skip_mut_forwarding_impls);
286     syn::custom_keyword!(features);
287     syn::custom_keyword!(require_store_data_send);
288     syn::custom_keyword!(wasmtime_crate);
289     syn::custom_keyword!(include_generated_code_from_file);
290 }
291 
292 enum Opt {
293     World(syn::LitStr),
294     Path(Vec<syn::LitStr>),
295     Inline(syn::LitStr),
296     Tracing(bool),
297     VerboseTracing(bool),
298     Async(AsyncConfig, Span),
299     TrappableErrorType(Vec<TrappableError>),
300     Ownership(Ownership),
301     Interfaces(syn::LitStr),
302     With(HashMap<String, String>),
303     TrappableImports(TrappableImports),
304     AdditionalDerives(Vec<syn::Path>),
305     Stringify(bool),
306     SkipMutForwardingImpls(bool),
307     Features(Vec<syn::LitStr>),
308     RequireStoreDataSend(bool),
309     WasmtimeCrate(syn::Path),
310     IncludeGeneratedCodeFromFile(bool),
311 }
312 
313 impl Parse for Opt {
314     fn parse(input: ParseStream<'_>) -> Result<Self> {
315         let l = input.lookahead1();
316         if l.peek(kw::path) {
317             input.parse::<kw::path>()?;
318             input.parse::<Token![:]>()?;
319 
320             let mut paths: Vec<syn::LitStr> = vec![];
321 
322             let l = input.lookahead1();
323             if l.peek(syn::LitStr) {
324                 paths.push(input.parse()?);
325             } else if l.peek(syn::token::Bracket) {
326                 let contents;
327                 syn::bracketed!(contents in input);
328                 let list = Punctuated::<_, Token![,]>::parse_terminated(&contents)?;
329 
330                 paths.extend(list.into_iter());
331             } else {
332                 return Err(l.error());
333             };
334 
335             Ok(Opt::Path(paths))
336         } else if l.peek(kw::inline) {
337             input.parse::<kw::inline>()?;
338             input.parse::<Token![:]>()?;
339             Ok(Opt::Inline(input.parse()?))
340         } else if l.peek(kw::world) {
341             input.parse::<kw::world>()?;
342             input.parse::<Token![:]>()?;
343             Ok(Opt::World(input.parse()?))
344         } else if l.peek(kw::tracing) {
345             input.parse::<kw::tracing>()?;
346             input.parse::<Token![:]>()?;
347             Ok(Opt::Tracing(input.parse::<syn::LitBool>()?.value))
348         } else if l.peek(kw::verbose_tracing) {
349             input.parse::<kw::verbose_tracing>()?;
350             input.parse::<Token![:]>()?;
351             Ok(Opt::VerboseTracing(input.parse::<syn::LitBool>()?.value))
352         } else if l.peek(Token![async]) {
353             let span = input.parse::<Token![async]>()?.span;
354             input.parse::<Token![:]>()?;
355             if input.peek(syn::LitBool) {
356                 match input.parse::<syn::LitBool>()?.value {
357                     true => Ok(Opt::Async(AsyncConfig::All, span)),
358                     false => Ok(Opt::Async(AsyncConfig::None, span)),
359                 }
360             } else {
361                 let contents;
362                 syn::braced!(contents in input);
363 
364                 let l = contents.lookahead1();
365                 let ctor: fn(HashSet<String>) -> AsyncConfig = if l.peek(kw::except_imports) {
366                     contents.parse::<kw::except_imports>()?;
367                     contents.parse::<Token![:]>()?;
368                     AsyncConfig::AllExceptImports
369                 } else if l.peek(kw::only_imports) {
370                     contents.parse::<kw::only_imports>()?;
371                     contents.parse::<Token![:]>()?;
372                     AsyncConfig::OnlyImports
373                 } else {
374                     return Err(l.error());
375                 };
376 
377                 let list;
378                 syn::bracketed!(list in contents);
379                 let fields: Punctuated<syn::LitStr, Token![,]> =
380                     list.parse_terminated(Parse::parse, Token![,])?;
381 
382                 if contents.peek(Token![,]) {
383                     contents.parse::<Token![,]>()?;
384                 }
385                 Ok(Opt::Async(
386                     ctor(fields.iter().map(|s| s.value()).collect()),
387                     span,
388                 ))
389             }
390         } else if l.peek(kw::ownership) {
391             input.parse::<kw::ownership>()?;
392             input.parse::<Token![:]>()?;
393             let ownership = input.parse::<syn::Ident>()?;
394             Ok(Opt::Ownership(match ownership.to_string().as_str() {
395                 "Owning" => Ownership::Owning,
396                 "Borrowing" => Ownership::Borrowing {
397                     duplicate_if_necessary: {
398                         let contents;
399                         braced!(contents in input);
400                         let field = contents.parse::<syn::Ident>()?;
401                         match field.to_string().as_str() {
402                             "duplicate_if_necessary" => {
403                                 contents.parse::<Token![:]>()?;
404                                 contents.parse::<syn::LitBool>()?.value
405                             }
406                             name => {
407                                 return Err(Error::new(
408                                     field.span(),
409                                     format!(
410                                         "unrecognized `Ownership::Borrowing` field: `{name}`; \
411                                          expected `duplicate_if_necessary`"
412                                     ),
413                                 ));
414                             }
415                         }
416                     },
417                 },
418                 name => {
419                     return Err(Error::new(
420                         ownership.span(),
421                         format!(
422                             "unrecognized ownership: `{name}`; \
423                              expected `Owning` or `Borrowing`"
424                         ),
425                     ));
426                 }
427             }))
428         } else if l.peek(kw::trappable_error_type) {
429             input.parse::<kw::trappable_error_type>()?;
430             input.parse::<Token![:]>()?;
431             let contents;
432             let _lbrace = braced!(contents in input);
433             let fields: Punctuated<_, Token![,]> =
434                 contents.parse_terminated(trappable_error_field_parse, Token![,])?;
435             Ok(Opt::TrappableErrorType(Vec::from_iter(fields)))
436         } else if l.peek(kw::interfaces) {
437             input.parse::<kw::interfaces>()?;
438             input.parse::<Token![:]>()?;
439             Ok(Opt::Interfaces(input.parse::<syn::LitStr>()?))
440         } else if l.peek(kw::with) {
441             input.parse::<kw::with>()?;
442             input.parse::<Token![:]>()?;
443             let contents;
444             let _lbrace = braced!(contents in input);
445             let fields: Punctuated<(String, String), Token![,]> =
446                 contents.parse_terminated(with_field_parse, Token![,])?;
447             Ok(Opt::With(HashMap::from_iter(fields)))
448         } else if l.peek(kw::trappable_imports) {
449             input.parse::<kw::trappable_imports>()?;
450             input.parse::<Token![:]>()?;
451             let config = if input.peek(syn::LitBool) {
452                 match input.parse::<syn::LitBool>()?.value {
453                     true => TrappableImports::All,
454                     false => TrappableImports::None,
455                 }
456             } else {
457                 let contents;
458                 syn::bracketed!(contents in input);
459                 let fields: Punctuated<syn::LitStr, Token![,]> =
460                     contents.parse_terminated(Parse::parse, Token![,])?;
461                 TrappableImports::Only(fields.iter().map(|s| s.value()).collect())
462             };
463             Ok(Opt::TrappableImports(config))
464         } else if l.peek(kw::additional_derives) {
465             input.parse::<kw::additional_derives>()?;
466             input.parse::<Token![:]>()?;
467             let contents;
468             syn::bracketed!(contents in input);
469             let list = Punctuated::<_, Token![,]>::parse_terminated(&contents)?;
470             Ok(Opt::AdditionalDerives(list.iter().cloned().collect()))
471         } else if l.peek(kw::stringify) {
472             input.parse::<kw::stringify>()?;
473             input.parse::<Token![:]>()?;
474             Ok(Opt::Stringify(input.parse::<syn::LitBool>()?.value))
475         } else if l.peek(kw::skip_mut_forwarding_impls) {
476             input.parse::<kw::skip_mut_forwarding_impls>()?;
477             input.parse::<Token![:]>()?;
478             Ok(Opt::SkipMutForwardingImpls(
479                 input.parse::<syn::LitBool>()?.value,
480             ))
481         } else if l.peek(kw::features) {
482             input.parse::<kw::features>()?;
483             input.parse::<Token![:]>()?;
484             let contents;
485             syn::bracketed!(contents in input);
486             let list = Punctuated::<_, Token![,]>::parse_terminated(&contents)?;
487             Ok(Opt::Features(list.into_iter().collect()))
488         } else if l.peek(kw::require_store_data_send) {
489             input.parse::<kw::require_store_data_send>()?;
490             input.parse::<Token![:]>()?;
491             Ok(Opt::RequireStoreDataSend(
492                 input.parse::<syn::LitBool>()?.value,
493             ))
494         } else if l.peek(kw::wasmtime_crate) {
495             input.parse::<kw::wasmtime_crate>()?;
496             input.parse::<Token![:]>()?;
497             Ok(Opt::WasmtimeCrate(input.parse()?))
498         } else if l.peek(kw::include_generated_code_from_file) {
499             input.parse::<kw::include_generated_code_from_file>()?;
500             input.parse::<Token![:]>()?;
501             Ok(Opt::IncludeGeneratedCodeFromFile(
502                 input.parse::<syn::LitBool>()?.value,
503             ))
504         } else {
505             Err(l.error())
506         }
507     }
508 }
509 
510 fn trappable_error_field_parse(input: ParseStream<'_>) -> Result<TrappableError> {
511     let wit_path = input.parse::<syn::LitStr>()?.value();
512     input.parse::<Token![=>]>()?;
513     let rust_type_name = input.parse::<syn::Path>()?.to_token_stream().to_string();
514     Ok(TrappableError {
515         wit_path,
516         rust_type_name,
517     })
518 }
519 
520 fn with_field_parse(input: ParseStream<'_>) -> Result<(String, String)> {
521     let interface = input.parse::<syn::LitStr>()?.value();
522     input.parse::<Token![:]>()?;
523     let start = input.span();
524     let path = input.parse::<syn::Path>()?;
525 
526     // It's not possible for the segments of a path to be empty
527     let span = start
528         .join(path.segments.last().unwrap().ident.span())
529         .unwrap_or(start);
530 
531     let mut buf = String::new();
532     let append = |buf: &mut String, segment: syn::PathSegment| -> Result<()> {
533         if segment.arguments != syn::PathArguments::None {
534             return Err(Error::new(
535                 span,
536                 "Module path must not contain angles or parens",
537             ));
538         }
539 
540         buf.push_str(&segment.ident.to_string());
541 
542         Ok(())
543     };
544 
545     if path.leading_colon.is_some() {
546         buf.push_str("::");
547     }
548 
549     let mut segments = path.segments.into_iter();
550 
551     if let Some(segment) = segments.next() {
552         append(&mut buf, segment)?;
553     }
554 
555     for segment in segments {
556         buf.push_str("::");
557         append(&mut buf, segment)?;
558     }
559 
560     Ok((interface, buf))
561 }
562