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