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 path = None;
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(s) => {
95                         if path.is_some() {
96                             return Err(Error::new(s.span(), "cannot specify second path"));
97                         }
98                         path = Some(s.value());
99                     }
100                     Opt::World(s) => {
101                         if world.is_some() {
102                             return Err(Error::new(s.span(), "cannot specify second world"));
103                         }
104                         world = Some(s.value());
105                     }
106                     Opt::Inline(s) => {
107                         if inline.is_some() {
108                             return Err(Error::new(s.span(), "cannot specify second source"));
109                         }
110                         inline = Some(s.value());
111                     }
112                     Opt::Tracing(val) => opts.tracing = val,
113                     Opt::Async(val, span) => {
114                         if async_configured {
115                             return Err(Error::new(span, "cannot specify second async config"));
116                         }
117                         async_configured = true;
118                         opts.async_ = val;
119                     }
120                     Opt::TrappableErrorType(val) => opts.trappable_error_type = val,
121                     Opt::TrappableImports(val) => opts.trappable_imports = val,
122                     Opt::Ownership(val) => opts.ownership = val,
123                     Opt::Interfaces(s) => {
124                         if inline.is_some() {
125                             return Err(Error::new(s.span(), "cannot specify a second source"));
126                         }
127                         inline = Some(format!(
128                             "
129                                 package wasmtime:component-macro-synthesized;
130 
131                                 world interfaces {{
132                                     {}
133                                 }}
134                             ",
135                             s.value()
136                         ));
137 
138                         if world.is_some() {
139                             return Err(Error::new(
140                                 s.span(),
141                                 "cannot specify a world with `interfaces`",
142                             ));
143                         }
144                         world = Some("interfaces".to_string());
145 
146                         opts.only_interfaces = true;
147                     }
148                     Opt::With(val) => opts.with.extend(val),
149                     Opt::AdditionalDerives(paths) => {
150                         opts.additional_derive_attributes = paths
151                             .into_iter()
152                             .map(|p| p.into_token_stream().to_string())
153                             .collect()
154                     }
155                     Opt::Stringify(val) => opts.stringify = val,
156                     Opt::SkipMutForwardingImpls(val) => opts.skip_mut_forwarding_impls = val,
157                     Opt::Features(f) => {
158                         features.extend(f.into_iter().map(|f| f.value()));
159                     }
160                     Opt::RequireStoreDataSend(val) => opts.require_store_data_send = val,
161                     Opt::WasmtimeCrate(f) => {
162                         opts.wasmtime_crate = Some(f.into_token_stream().to_string())
163                     }
164                     Opt::IncludeGeneratedCodeFromFile(i) => include_generated_code_from_file = i,
165                 }
166             }
167         } else {
168             world = input.parse::<Option<syn::LitStr>>()?.map(|s| s.value());
169             if input.parse::<Option<syn::token::In>>()?.is_some() {
170                 path = Some(input.parse::<syn::LitStr>()?.value());
171             }
172         }
173         let (resolve, pkg, files) = parse_source(&path, &inline, &features)
174             .map_err(|err| Error::new(call_site, format!("{err:?}")))?;
175 
176         let world = resolve
177             .select_world(pkg, world.as_deref())
178             .map_err(|e| Error::new(call_site, format!("{e:?}")))?;
179         Ok(Config {
180             opts,
181             resolve,
182             world,
183             files,
184             include_generated_code_from_file,
185         })
186     }
187 }
188 
189 fn parse_source(
190     path: &Option<String>,
191     inline: &Option<String>,
192     features: &[String],
193 ) -> anyhow::Result<(Resolve, PackageId, Vec<PathBuf>)> {
194     let mut resolve = Resolve::default();
195     resolve.features.extend(features.iter().cloned());
196     let mut files = Vec::new();
197     let root = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap());
198 
199     let mut parse = |resolve: &mut Resolve, path: &Path| -> anyhow::Result<_> {
200         // Try to normalize the path to make the error message more understandable when
201         // the path is not correct. Fallback to the original path if normalization fails
202         // (probably return an error somewhere else).
203         let normalized_path = match std::fs::canonicalize(path) {
204             Ok(p) => p,
205             Err(_) => path.to_path_buf(),
206         };
207         let (pkg, sources) = resolve.push_path(normalized_path)?;
208         files.extend(sources);
209         Ok(pkg)
210     };
211 
212     let path_pkg = if let Some(path) = path {
213         Some(parse(&mut resolve, &root.join(path))?)
214     } else {
215         None
216     };
217 
218     let inline_pkgs = if let Some(inline) = inline {
219         Some(resolve.push_group(UnresolvedPackageGroup::parse("macro-input", inline)?)?)
220     } else {
221         None
222     };
223 
224     let pkgs = inline_pkgs
225         .or(path_pkg)
226         .map_or_else(|| parse(&mut resolve, &root.join("wit")), Ok)?;
227 
228     Ok((resolve, pkgs, files))
229 }
230 
231 mod kw {
232     syn::custom_keyword!(inline);
233     syn::custom_keyword!(path);
234     syn::custom_keyword!(tracing);
235     syn::custom_keyword!(trappable_error_type);
236     syn::custom_keyword!(world);
237     syn::custom_keyword!(ownership);
238     syn::custom_keyword!(interfaces);
239     syn::custom_keyword!(with);
240     syn::custom_keyword!(except_imports);
241     syn::custom_keyword!(only_imports);
242     syn::custom_keyword!(trappable_imports);
243     syn::custom_keyword!(additional_derives);
244     syn::custom_keyword!(stringify);
245     syn::custom_keyword!(skip_mut_forwarding_impls);
246     syn::custom_keyword!(features);
247     syn::custom_keyword!(require_store_data_send);
248     syn::custom_keyword!(wasmtime_crate);
249     syn::custom_keyword!(include_generated_code_from_file);
250 }
251 
252 enum Opt {
253     World(syn::LitStr),
254     Path(syn::LitStr),
255     Inline(syn::LitStr),
256     Tracing(bool),
257     Async(AsyncConfig, Span),
258     TrappableErrorType(Vec<TrappableError>),
259     Ownership(Ownership),
260     Interfaces(syn::LitStr),
261     With(HashMap<String, String>),
262     TrappableImports(TrappableImports),
263     AdditionalDerives(Vec<syn::Path>),
264     Stringify(bool),
265     SkipMutForwardingImpls(bool),
266     Features(Vec<syn::LitStr>),
267     RequireStoreDataSend(bool),
268     WasmtimeCrate(syn::Path),
269     IncludeGeneratedCodeFromFile(bool),
270 }
271 
272 impl Parse for Opt {
273     fn parse(input: ParseStream<'_>) -> Result<Self> {
274         let l = input.lookahead1();
275         if l.peek(kw::path) {
276             input.parse::<kw::path>()?;
277             input.parse::<Token![:]>()?;
278             Ok(Opt::Path(input.parse()?))
279         } else if l.peek(kw::inline) {
280             input.parse::<kw::inline>()?;
281             input.parse::<Token![:]>()?;
282             Ok(Opt::Inline(input.parse()?))
283         } else if l.peek(kw::world) {
284             input.parse::<kw::world>()?;
285             input.parse::<Token![:]>()?;
286             Ok(Opt::World(input.parse()?))
287         } else if l.peek(kw::tracing) {
288             input.parse::<kw::tracing>()?;
289             input.parse::<Token![:]>()?;
290             Ok(Opt::Tracing(input.parse::<syn::LitBool>()?.value))
291         } else if l.peek(Token![async]) {
292             let span = input.parse::<Token![async]>()?.span;
293             input.parse::<Token![:]>()?;
294             if input.peek(syn::LitBool) {
295                 match input.parse::<syn::LitBool>()?.value {
296                     true => Ok(Opt::Async(AsyncConfig::All, span)),
297                     false => Ok(Opt::Async(AsyncConfig::None, span)),
298                 }
299             } else {
300                 let contents;
301                 syn::braced!(contents in input);
302 
303                 let l = contents.lookahead1();
304                 let ctor: fn(HashSet<String>) -> AsyncConfig = if l.peek(kw::except_imports) {
305                     contents.parse::<kw::except_imports>()?;
306                     contents.parse::<Token![:]>()?;
307                     AsyncConfig::AllExceptImports
308                 } else if l.peek(kw::only_imports) {
309                     contents.parse::<kw::only_imports>()?;
310                     contents.parse::<Token![:]>()?;
311                     AsyncConfig::OnlyImports
312                 } else {
313                     return Err(l.error());
314                 };
315 
316                 let list;
317                 syn::bracketed!(list in contents);
318                 let fields: Punctuated<syn::LitStr, Token![,]> =
319                     list.parse_terminated(Parse::parse, Token![,])?;
320 
321                 if contents.peek(Token![,]) {
322                     contents.parse::<Token![,]>()?;
323                 }
324                 Ok(Opt::Async(
325                     ctor(fields.iter().map(|s| s.value()).collect()),
326                     span,
327                 ))
328             }
329         } else if l.peek(kw::ownership) {
330             input.parse::<kw::ownership>()?;
331             input.parse::<Token![:]>()?;
332             let ownership = input.parse::<syn::Ident>()?;
333             Ok(Opt::Ownership(match ownership.to_string().as_str() {
334                 "Owning" => Ownership::Owning,
335                 "Borrowing" => Ownership::Borrowing {
336                     duplicate_if_necessary: {
337                         let contents;
338                         braced!(contents in input);
339                         let field = contents.parse::<syn::Ident>()?;
340                         match field.to_string().as_str() {
341                             "duplicate_if_necessary" => {
342                                 contents.parse::<Token![:]>()?;
343                                 contents.parse::<syn::LitBool>()?.value
344                             }
345                             name => {
346                                 return Err(Error::new(
347                                     field.span(),
348                                     format!(
349                                         "unrecognized `Ownership::Borrowing` field: `{name}`; \
350                                          expected `duplicate_if_necessary`"
351                                     ),
352                                 ));
353                             }
354                         }
355                     },
356                 },
357                 name => {
358                     return Err(Error::new(
359                         ownership.span(),
360                         format!(
361                             "unrecognized ownership: `{name}`; \
362                              expected `Owning` or `Borrowing`"
363                         ),
364                     ));
365                 }
366             }))
367         } else if l.peek(kw::trappable_error_type) {
368             input.parse::<kw::trappable_error_type>()?;
369             input.parse::<Token![:]>()?;
370             let contents;
371             let _lbrace = braced!(contents in input);
372             let fields: Punctuated<_, Token![,]> =
373                 contents.parse_terminated(trappable_error_field_parse, Token![,])?;
374             Ok(Opt::TrappableErrorType(Vec::from_iter(fields)))
375         } else if l.peek(kw::interfaces) {
376             input.parse::<kw::interfaces>()?;
377             input.parse::<Token![:]>()?;
378             Ok(Opt::Interfaces(input.parse::<syn::LitStr>()?))
379         } else if l.peek(kw::with) {
380             input.parse::<kw::with>()?;
381             input.parse::<Token![:]>()?;
382             let contents;
383             let _lbrace = braced!(contents in input);
384             let fields: Punctuated<(String, String), Token![,]> =
385                 contents.parse_terminated(with_field_parse, Token![,])?;
386             Ok(Opt::With(HashMap::from_iter(fields)))
387         } else if l.peek(kw::trappable_imports) {
388             input.parse::<kw::trappable_imports>()?;
389             input.parse::<Token![:]>()?;
390             let config = if input.peek(syn::LitBool) {
391                 match input.parse::<syn::LitBool>()?.value {
392                     true => TrappableImports::All,
393                     false => TrappableImports::None,
394                 }
395             } else {
396                 let contents;
397                 syn::bracketed!(contents in input);
398                 let fields: Punctuated<syn::LitStr, Token![,]> =
399                     contents.parse_terminated(Parse::parse, Token![,])?;
400                 TrappableImports::Only(fields.iter().map(|s| s.value()).collect())
401             };
402             Ok(Opt::TrappableImports(config))
403         } else if l.peek(kw::additional_derives) {
404             input.parse::<kw::additional_derives>()?;
405             input.parse::<Token![:]>()?;
406             let contents;
407             syn::bracketed!(contents in input);
408             let list = Punctuated::<_, Token![,]>::parse_terminated(&contents)?;
409             Ok(Opt::AdditionalDerives(list.iter().cloned().collect()))
410         } else if l.peek(kw::stringify) {
411             input.parse::<kw::stringify>()?;
412             input.parse::<Token![:]>()?;
413             Ok(Opt::Stringify(input.parse::<syn::LitBool>()?.value))
414         } else if l.peek(kw::skip_mut_forwarding_impls) {
415             input.parse::<kw::skip_mut_forwarding_impls>()?;
416             input.parse::<Token![:]>()?;
417             Ok(Opt::SkipMutForwardingImpls(
418                 input.parse::<syn::LitBool>()?.value,
419             ))
420         } else if l.peek(kw::features) {
421             input.parse::<kw::features>()?;
422             input.parse::<Token![:]>()?;
423             let contents;
424             syn::bracketed!(contents in input);
425             let list = Punctuated::<_, Token![,]>::parse_terminated(&contents)?;
426             Ok(Opt::Features(list.into_iter().collect()))
427         } else if l.peek(kw::require_store_data_send) {
428             input.parse::<kw::require_store_data_send>()?;
429             input.parse::<Token![:]>()?;
430             Ok(Opt::RequireStoreDataSend(
431                 input.parse::<syn::LitBool>()?.value,
432             ))
433         } else if l.peek(kw::wasmtime_crate) {
434             input.parse::<kw::wasmtime_crate>()?;
435             input.parse::<Token![:]>()?;
436             Ok(Opt::WasmtimeCrate(input.parse()?))
437         } else if l.peek(kw::include_generated_code_from_file) {
438             input.parse::<kw::include_generated_code_from_file>()?;
439             input.parse::<Token![:]>()?;
440             Ok(Opt::IncludeGeneratedCodeFromFile(
441                 input.parse::<syn::LitBool>()?.value,
442             ))
443         } else {
444             Err(l.error())
445         }
446     }
447 }
448 
449 fn trappable_error_field_parse(input: ParseStream<'_>) -> Result<TrappableError> {
450     let wit_path = input.parse::<syn::LitStr>()?.value();
451     input.parse::<Token![=>]>()?;
452     let rust_type_name = input.parse::<syn::Path>()?.to_token_stream().to_string();
453     Ok(TrappableError {
454         wit_path,
455         rust_type_name,
456     })
457 }
458 
459 fn with_field_parse(input: ParseStream<'_>) -> Result<(String, String)> {
460     let interface = input.parse::<syn::LitStr>()?.value();
461     input.parse::<Token![:]>()?;
462     let start = input.span();
463     let path = input.parse::<syn::Path>()?;
464 
465     // It's not possible for the segments of a path to be empty
466     let span = start
467         .join(path.segments.last().unwrap().ident.span())
468         .unwrap_or(start);
469 
470     let mut buf = String::new();
471     let append = |buf: &mut String, segment: syn::PathSegment| -> Result<()> {
472         if segment.arguments != syn::PathArguments::None {
473             return Err(Error::new(
474                 span,
475                 "Module path must not contain angles or parens",
476             ));
477         }
478 
479         buf.push_str(&segment.ident.to_string());
480 
481         Ok(())
482     };
483 
484     if path.leading_colon.is_some() {
485         buf.push_str("::");
486     }
487 
488     let mut segments = path.segments.into_iter();
489 
490     if let Some(segment) = segments.next() {
491         append(&mut buf, segment)?;
492     }
493 
494     for segment in segments {
495         buf.push_str("::");
496         append(&mut buf, segment)?;
497     }
498 
499     Ok((interface, buf))
500 }
501