use std::{collections::HashSet, env, fs, path::PathBuf};
use proc_macro2::{Span, TokenStream as TokenStream2};
use quote::{ToTokens, quote};
use syn::{
Error, FnArg, Item, ItemFn, LitStr, Pat, Token, TypePath,
parse::{Parse, ParseStream},
spanned::Spanned,
visit_mut::VisitMut,
};
use wit_bindgen_core::{
WorldGenerator,
wit_parser::{
Function, Handle, InterfaceId, PackageId, Resolve, Type as WitType, TypeDefKind, TypeId,
TypeOwner, UnresolvedPackageGroup, WorldId, WorldItem,
},
};
use wit_bindgen_rust::{Opts, WithOption};
use crate::{fpi, manifest_paths};
const CORE_TYPES_INTERFACE: &str = "miden:base/core-types@1.0.0";
#[derive(Default)]
struct GenerateArgs {
inline: Option<LitStr>,
with_entries: Vec<(String, WithOption)>,
}
fn parse_with_entry(input: ParseStream<'_>) -> syn::Result<(String, WithOption)> {
let key: LitStr = input.parse()?;
input.parse::<Token![:]>()?;
let path: syn::Path = input.parse()?;
let option = if path.leading_colon.is_none()
&& path.segments.len() == 1
&& path.segments.first().is_some_and(|seg| seg.ident == "generate")
{
WithOption::Generate
} else {
let path_str = path.to_token_stream().to_string().replace(' ', "");
WithOption::Path(path_str)
};
Ok((key.value(), option))
}
impl Parse for GenerateArgs {
fn parse(input: ParseStream<'_>) -> syn::Result<Self> {
let mut args = GenerateArgs::default();
while !input.is_empty() {
let ident: syn::Ident = input.parse()?;
let name = ident.to_string();
input.parse::<Token![=]>()?;
if name == "inline" {
if args.inline.is_some() {
return Err(syn::Error::new(ident.span(), "duplicate `inline` argument"));
}
args.inline = Some(input.parse()?);
} else if name == "with" {
if !args.with_entries.is_empty() {
return Err(syn::Error::new(ident.span(), "duplicate `with` argument"));
}
let content;
syn::braced!(content in input);
while !content.is_empty() {
args.with_entries.push(parse_with_entry(&content)?);
if content.peek(Token![,]) {
content.parse::<Token![,]>()?;
}
}
} else {
return Err(syn::Error::new(
ident.span(),
format!("unsupported generate! argument `{name}`"),
));
}
if input.peek(Token![,]) {
let _ = input.parse::<Token![,]>()?;
}
}
Ok(args)
}
}
pub(crate) fn expand(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
let input_tokens: proc_macro2::TokenStream = input.into();
let args = if input_tokens.is_empty() {
GenerateArgs::default()
} else {
match syn::parse2::<GenerateArgs>(input_tokens) {
Ok(parsed) => parsed,
Err(err) => return err.to_compile_error().into(),
}
};
let resolve_opts = manifest_paths::ResolveOptions {
allow_missing_local_wit: args.inline.is_some(),
};
match manifest_paths::resolve_wit_paths(resolve_opts) {
Ok(config) => {
if config.paths.is_empty() {
return Error::new(
Span::call_site(),
"no WIT dependencies declared under \
[package.metadata.component.target.dependencies]",
)
.to_compile_error()
.into();
}
let inline_world = args
.inline
.as_ref()
.and_then(|src| manifest_paths::extract_world_name(&src.value()));
let world_value = inline_world.or_else(|| config.world.clone());
if args.inline.is_some() && world_value.is_none() {
return Error::new(
Span::call_site(),
"failed to detect world name for inline WIT provided to generate!",
)
.to_compile_error()
.into();
}
match generate_bindings(&args, &config, world_value.as_deref()) {
Ok(raw_bindings) => quote! {
#[doc(hidden)]
#[allow(dead_code)]
pub mod bindings {
#raw_bindings
}
}
.into(),
Err(err) => err.to_compile_error().into(),
}
}
Err(err) => err.to_compile_error().into(),
}
}
fn generate_bindings(
args: &GenerateArgs,
config: &manifest_paths::ResolvedWit,
world: Option<&str>,
) -> Result<TokenStream2, Error> {
generate_bindings_from_sources(
&config.paths,
args.inline.as_ref().map(|src| src.value()).as_deref(),
world,
&args.with_entries,
&[],
)
}
pub(crate) fn generate_inline_fpi_bindings(
config: &manifest_paths::ResolvedWit,
inline_source: &str,
world: &str,
fpi_imports: &[String],
with_entries: &[(String, WithOption)],
) -> Result<TokenStream2, Error> {
generate_bindings_from_sources(
&config.paths,
Some(inline_source),
Some(world),
with_entries,
fpi_imports,
)
}
pub(crate) fn generate_inline_import_bindings(
config: &manifest_paths::ResolvedWit,
inline_source: &str,
world: &str,
with_entries: &[(String, WithOption)],
) -> Result<TokenStream2, Error> {
generate_bindings_from_sources(
&config.paths,
Some(inline_source),
Some(world),
with_entries,
&[],
)
}
fn generate_bindings_from_sources(
paths: &[String],
inline_source: Option<&str>,
world: Option<&str>,
with_entries: &[(String, WithOption)],
fpi_imports: &[String],
) -> Result<TokenStream2, Error> {
let mut wit_sources = load_wit_sources(paths, inline_source)?;
let world_id = wit_sources
.resolve
.select_world(&wit_sources.packages, world)
.map_err(|err| Error::new(Span::call_site(), err.to_string()))?;
fpi::inject_imports(&mut wit_sources.resolve, world_id, fpi_imports)?;
let mut opts = Opts {
generate_all: true,
runtime_path: Some("::miden::wit_bindgen::rt".to_string()),
default_bindings_module: Some("bindings".to_string()),
..Opts::default()
};
push_custom_with_entries(&mut opts, with_entries);
if world_uses_miden_core_types(&wit_sources.resolve, world_id) {
push_default_with_entries(&mut opts);
}
let mut generated_files = wit_bindgen_core::Files::default();
let mut generator = opts.build();
generator
.generate(&mut wit_sources.resolve, world_id, &mut generated_files)
.map_err(|err| Error::new(Span::call_site(), err.to_string()))?;
let (_, src_bytes) = generated_files
.iter()
.next()
.ok_or_else(|| Error::new(Span::call_site(), "wit-bindgen emitted no bindings"))?;
let src = std::str::from_utf8(src_bytes)
.map_err(|err| Error::new(Span::call_site(), format!("invalid UTF-8: {err}")))?;
let mut tokens: TokenStream2 = src
.parse()
.map_err(|err| Error::new(Span::call_site(), format!("failed to parse bindings: {err}")))?;
for path in wit_sources.files_read {
let utf8_path = path.to_str().ok_or_else(|| {
Error::new(
Span::call_site(),
format!("path '{}' contains invalid UTF-8", path.display()),
)
})?;
tokens.extend(quote! {
const _: &[u8] = include_bytes!(#utf8_path);
});
}
Ok(tokens)
}
struct LoadedWitSources {
resolve: Resolve,
packages: Vec<PackageId>,
files_read: Vec<PathBuf>,
}
fn load_wit_sources(
paths: &[String],
inline_source: Option<&str>,
) -> Result<LoadedWitSources, Error> {
let manifest_dir = env::var("CARGO_MANIFEST_DIR").map_err(|err| {
Error::new(Span::call_site(), format!("failed to read CARGO_MANIFEST_DIR: {err}"))
})?;
let manifest_dir = PathBuf::from(manifest_dir);
let mut resolve = Resolve::default();
let mut packages = Vec::new();
let mut files = Vec::new();
for path in paths {
let path_buf = PathBuf::from(path);
let absolute = if path_buf.is_absolute() {
path_buf
} else {
manifest_dir.join(path_buf)
};
let normalized = fs::canonicalize(&absolute).unwrap_or(absolute);
let (pkg, sources) = resolve.push_path(normalized.clone()).map_err(|err| {
Error::new(
Span::call_site(),
format!("failed to load WIT from '{}': {err}", normalized.display()),
)
})?;
packages.push(pkg);
files.extend(sources.paths().map(|p| p.to_owned()));
}
if let Some(src) = inline_source {
packages.clear();
let group = UnresolvedPackageGroup::parse("inline", src)
.map_err(|err| Error::new(Span::call_site(), err.to_string()))?;
let pkg = resolve
.push_group(group)
.map_err(|err| Error::new(Span::call_site(), err.to_string()))?;
packages.push(pkg);
}
Ok(LoadedWitSources {
resolve,
packages,
files_read: files,
})
}
fn push_custom_with_entries(opts: &mut Opts, entries: &[(String, WithOption)]) {
opts.with.extend(entries.iter().cloned());
}
fn push_default_with_entries(opts: &mut Opts) {
opts.with.push((CORE_TYPES_INTERFACE.to_string(), WithOption::Generate));
push_path_entry(opts, &format!("{CORE_TYPES_INTERFACE}/felt"), "::miden::Felt");
push_path_entry(opts, &format!("{CORE_TYPES_INTERFACE}/word"), "::miden::Word");
push_path_entry(opts, &format!("{CORE_TYPES_INTERFACE}/asset"), "::miden::Asset");
push_path_entry(opts, &format!("{CORE_TYPES_INTERFACE}/account-id"), "::miden::AccountId");
push_path_entry(opts, &format!("{CORE_TYPES_INTERFACE}/tag"), "::miden::Tag");
push_path_entry(opts, &format!("{CORE_TYPES_INTERFACE}/note-type"), "::miden::NoteType");
push_path_entry(opts, &format!("{CORE_TYPES_INTERFACE}/recipient"), "::miden::Recipient");
push_path_entry(opts, &format!("{CORE_TYPES_INTERFACE}/note-idx"), "::miden::NoteIdx");
}
fn push_path_entry(opts: &mut Opts, key: &str, value: &str) {
opts.with.push((key.to_string(), WithOption::Path(value.to_string())));
}
fn world_uses_miden_core_types(resolve: &Resolve, world_id: WorldId) -> bool {
let world = &resolve.worlds[world_id];
world
.imports
.values()
.chain(world.exports.values())
.any(|item| world_item_uses_interface(resolve, item, CORE_TYPES_INTERFACE))
}
fn world_item_uses_interface(resolve: &Resolve, item: &WorldItem, interface_path: &str) -> bool {
match item {
WorldItem::Interface { id, .. } => interface_uses_interface(resolve, *id, interface_path),
WorldItem::Function(function) => function_uses_interface(resolve, function, interface_path),
WorldItem::Type { id, .. } => {
type_id_uses_interface(resolve, *id, interface_path, &mut HashSet::new())
}
}
}
fn interface_uses_interface(
resolve: &Resolve,
interface_id: InterfaceId,
interface_path: &str,
) -> bool {
let interface = &resolve.interfaces[interface_id];
if interface_matches(resolve, interface_id, interface_path) {
return true;
}
let mut visited = HashSet::new();
interface.functions.values().any(|function| {
function_uses_interface_with_visited(resolve, function, interface_path, &mut visited)
}) || interface
.types
.values()
.any(|id| type_id_uses_interface(resolve, *id, interface_path, &mut visited))
}
fn function_uses_interface(resolve: &Resolve, function: &Function, interface_path: &str) -> bool {
function_uses_interface_with_visited(resolve, function, interface_path, &mut HashSet::new())
}
fn function_uses_interface_with_visited(
resolve: &Resolve,
function: &Function,
interface_path: &str,
visited: &mut HashSet<TypeId>,
) -> bool {
function
.params
.iter()
.any(|param| type_uses_interface(resolve, ¶m.ty, interface_path, visited))
|| function
.result
.as_ref()
.is_some_and(|ty| type_uses_interface(resolve, ty, interface_path, visited))
}
fn type_uses_interface(
resolve: &Resolve,
ty: &WitType,
interface_path: &str,
visited: &mut HashSet<TypeId>,
) -> bool {
match ty {
WitType::Id(id) => type_id_uses_interface(resolve, *id, interface_path, visited),
_ => false,
}
}
fn type_id_uses_interface(
resolve: &Resolve,
type_id: TypeId,
interface_path: &str,
visited: &mut HashSet<TypeId>,
) -> bool {
if !visited.insert(type_id) {
return false;
}
let def = &resolve.types[type_id];
if type_owner_uses_interface(resolve, def.owner, interface_path) {
return true;
}
match &def.kind {
TypeDefKind::Record(record) => record
.fields
.iter()
.any(|field| type_uses_interface(resolve, &field.ty, interface_path, visited)),
TypeDefKind::Tuple(tuple) => tuple
.types
.iter()
.any(|ty| type_uses_interface(resolve, ty, interface_path, visited)),
TypeDefKind::Variant(variant) => variant
.cases
.iter()
.filter_map(|case| case.ty.as_ref())
.any(|ty| type_uses_interface(resolve, ty, interface_path, visited)),
TypeDefKind::Option(ty)
| TypeDefKind::List(ty)
| TypeDefKind::FixedLengthList(ty, _)
| TypeDefKind::Type(ty)
| TypeDefKind::Future(Some(ty))
| TypeDefKind::Stream(Some(ty)) => {
type_uses_interface(resolve, ty, interface_path, visited)
}
TypeDefKind::Result(result) => result
.ok
.as_ref()
.into_iter()
.chain(result.err.as_ref())
.any(|ty| type_uses_interface(resolve, ty, interface_path, visited)),
TypeDefKind::Map(key, value) => {
type_uses_interface(resolve, key, interface_path, visited)
|| type_uses_interface(resolve, value, interface_path, visited)
}
TypeDefKind::Handle(Handle::Own(id) | Handle::Borrow(id)) => {
type_id_uses_interface(resolve, *id, interface_path, visited)
}
TypeDefKind::Resource
| TypeDefKind::Flags(_)
| TypeDefKind::Enum(_)
| TypeDefKind::Future(None)
| TypeDefKind::Stream(None)
| TypeDefKind::Unknown => false,
}
}
fn type_owner_uses_interface(resolve: &Resolve, owner: TypeOwner, interface_path: &str) -> bool {
match owner {
TypeOwner::World(_) => false,
TypeOwner::Interface(id) => interface_matches(resolve, id, interface_path),
TypeOwner::None => false,
}
}
fn interface_matches(resolve: &Resolve, interface_id: InterfaceId, interface_path: &str) -> bool {
let interface = &resolve.interfaces[interface_id];
let (Some(package_id), Some(name)) = (interface.package, interface.name.as_deref()) else {
return false;
};
resolve.packages[package_id].name.interface_id(name) == interface_path
}
pub(crate) fn qualify_signature_types(sig: &mut syn::Signature, module_path: &[syn::Ident]) {
struct TypeQualifier<'a> {
module_path: &'a [syn::Ident],
}
impl VisitMut for TypeQualifier<'_> {
fn visit_type_path_mut(&mut self, type_path: &mut TypePath) {
if type_path.qself.is_none()
&& type_path.path.leading_colon.is_none()
&& type_path.path.segments.len() == 1
{
let first_segment = &type_path.path.segments[0].ident;
let name = first_segment.to_string();
if is_primitive_or_std_type(&name) {
syn::visit_mut::visit_type_path_mut(self, type_path);
return;
}
let mut new_segments = syn::punctuated::Punctuated::new();
for ident in self.module_path {
new_segments.push(syn::PathSegment {
ident: ident.clone(),
arguments: syn::PathArguments::None,
});
}
new_segments.push(type_path.path.segments[0].clone());
type_path.path.segments = new_segments;
}
syn::visit_mut::visit_type_path_mut(self, type_path);
}
}
let mut qualifier = TypeQualifier { module_path };
qualifier.visit_signature_mut(sig);
}
fn is_primitive_or_std_type(name: &str) -> bool {
matches!(
name,
"bool"
| "char"
| "str"
| "u8"
| "u16"
| "u32"
| "u64"
| "u128"
| "usize"
| "i8"
| "i16"
| "i32"
| "i64"
| "i128"
| "isize"
| "f32"
| "f64"
| "String"
| "Vec"
| "Option"
| "Result"
| "Self"
)
}
pub(crate) fn collect_arg_idents(func: &ItemFn) -> syn::Result<Vec<syn::Ident>> {
func.sig
.inputs
.iter()
.map(|arg| match arg {
FnArg::Receiver(_) => {
Err(Error::new(func.sig.ident.span(), "unexpected receiver in generated function"))
}
FnArg::Typed(pat_type) => match pat_type.pat.as_ref() {
Pat::Ident(pat_ident) => Ok(pat_ident.ident.clone()),
other => Err(Error::new(
other.span(),
format!(
"unsupported argument pattern `{}` in generated function",
quote!(#other)
),
)),
},
})
.collect()
}
pub(crate) fn should_generate_struct(path: &[syn::Ident], items: &[Item]) -> bool {
if path.is_empty() {
return false;
}
let first = path[0].to_string();
if first == "exports" {
return false;
}
if first.starts_with('_') {
return false;
}
let last = path.last().unwrap().to_string();
if last.starts_with('_') {
return false;
}
!items.iter().any(|item| matches!(item, Item::Mod(_)))
}
pub(crate) fn format_module_path(path: &[syn::Ident]) -> String {
path.iter().map(|ident| ident.to_string()).collect::<Vec<_>>().join("::")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_should_generate_struct_empty_path() {
let empty_items: Vec<Item> = vec![];
assert!(!should_generate_struct(&[], &empty_items));
}
#[test]
fn test_should_generate_struct_exports_excluded() {
let empty_items: Vec<Item> = vec![];
let path = vec![syn::Ident::new("exports", Span::call_site())];
assert!(!should_generate_struct(&path, &empty_items));
let path = vec![
syn::Ident::new("exports", Span::call_site()),
syn::Ident::new("foo", Span::call_site()),
];
assert!(!should_generate_struct(&path, &empty_items));
}
#[test]
fn test_should_generate_struct_underscore_excluded() {
let empty_items: Vec<Item> = vec![];
let path = vec![syn::Ident::new("_private", Span::call_site())];
assert!(!should_generate_struct(&path, &empty_items));
let path = vec![
syn::Ident::new("miden", Span::call_site()),
syn::Ident::new("_internal", Span::call_site()),
];
assert!(!should_generate_struct(&path, &empty_items));
}
#[test]
fn test_should_generate_struct_valid_leaf_modules() {
let empty_items: Vec<Item> = vec![];
let path = vec![syn::Ident::new("miden", Span::call_site())];
assert!(should_generate_struct(&path, &empty_items));
let path = vec![
syn::Ident::new("miden", Span::call_site()),
syn::Ident::new("basic_wallet", Span::call_site()),
];
assert!(should_generate_struct(&path, &empty_items));
}
#[test]
fn test_should_generate_struct_non_leaf_excluded() {
let path = vec![syn::Ident::new("miden", Span::call_site())];
let items_with_mod: Vec<Item> = vec![syn::parse_quote! { mod nested {} }];
assert!(!should_generate_struct(&path, &items_with_mod));
let items_with_fn: Vec<Item> = vec![syn::parse_quote! { pub fn foo() {} }];
assert!(should_generate_struct(&path, &items_with_fn));
}
#[test]
fn test_format_module_path() {
let path = vec![
syn::Ident::new("miden", Span::call_site()),
syn::Ident::new("basic_wallet", Span::call_site()),
];
assert_eq!(format_module_path(&path), "miden::basic_wallet");
}
#[test]
fn test_format_module_path_empty() {
assert_eq!(format_module_path(&[]), "");
}
#[test]
fn test_collect_arg_idents() {
let func: ItemFn = syn::parse_quote! {
pub fn foo(a: u32, b: String, c: Vec<u8>) {}
};
let idents = collect_arg_idents(&func).unwrap();
let names: Vec<_> = idents.iter().map(|i| i.to_string()).collect();
assert_eq!(names, vec!["a", "b", "c"]);
}
#[test]
fn test_collect_arg_idents_empty() {
let func: ItemFn = syn::parse_quote! {
pub fn no_args() {}
};
let idents = collect_arg_idents(&func).unwrap();
assert!(idents.is_empty());
}
#[test]
fn test_qualify_signature_types() {
let mut sig: syn::Signature = syn::parse_quote! {
fn test_fn(a: StructA, b: u64) -> StructB
};
let path = vec![
syn::Ident::new("miden", Span::call_site()),
syn::Ident::new("component", Span::call_site()),
];
qualify_signature_types(&mut sig, &path);
let sig_str = sig.to_token_stream().to_string();
assert!(sig_str.contains("miden :: component :: StructA"));
assert!(sig_str.contains("miden :: component :: StructB"));
assert!(sig_str.contains("u64"));
assert!(!sig_str.contains("miden :: component :: u64"));
}
#[test]
fn test_qualify_signature_types_inside_option() {
let mut sig: syn::Signature = syn::parse_quote! {
fn roundtrip(payload: Option<OptionPayload>) -> Option<OptionPayload>
};
let path = vec![
syn::Ident::new("miden", Span::call_site()),
syn::Ident::new("account", Span::call_site()),
syn::Ident::new("interface", Span::call_site()),
];
qualify_signature_types(&mut sig, &path);
let signature = sig.to_token_stream().to_string().replace(' ', "");
assert!(signature.contains("payload:Option<miden::account::interface::OptionPayload>"));
assert!(signature.contains("->Option<miden::account::interface::OptionPayload>"));
}
#[test]
fn test_parse_with_entry_generate() {
let input: TokenStream2 = quote! { "miden:foo/bar": generate };
let parsed = syn::parse2::<GenerateArgs>(quote! { with = { #input } }).unwrap();
assert_eq!(parsed.with_entries.len(), 1);
assert_eq!(parsed.with_entries[0].0, "miden:foo/bar");
assert!(matches!(parsed.with_entries[0].1, WithOption::Generate));
}
#[test]
fn test_parse_with_entry_path() {
let input: TokenStream2 = quote! { "miden:foo/bar": ::my::custom::Type };
let parsed = syn::parse2::<GenerateArgs>(quote! { with = { #input } }).unwrap();
assert_eq!(parsed.with_entries.len(), 1);
assert_eq!(parsed.with_entries[0].0, "miden:foo/bar");
match &parsed.with_entries[0].1 {
WithOption::Path(p) => assert_eq!(p, "::my::custom::Type"),
_ => panic!("expected Path variant"),
}
}
#[test]
fn test_parse_multiple_with_entries() {
let parsed = syn::parse2::<GenerateArgs>(quote! {
with = {
"miden:a/b": generate,
"miden:c/d": ::foo::Bar
}
})
.unwrap();
assert_eq!(parsed.with_entries.len(), 2);
assert_eq!(parsed.with_entries[0].0, "miden:a/b");
assert_eq!(parsed.with_entries[1].0, "miden:c/d");
}
fn parse_test_world(source: &str) -> (Resolve, WorldId) {
let mut resolve = Resolve::default();
let sdk_group =
UnresolvedPackageGroup::parse("miden.wit", manifest_paths::SDK_WIT_SOURCE).unwrap();
resolve.push_group(sdk_group).unwrap();
let group = UnresolvedPackageGroup::parse("inline", source).unwrap();
let package = resolve.push_group(group).unwrap();
let world = resolve.select_world(&[package], None).unwrap();
(resolve, world)
}
#[test]
fn test_world_uses_miden_core_types_rejects_primitive_only_world() {
let (resolve, world) = parse_test_world(
r#"
package miden:primitive-variant@0.1.0;
interface primitive-variant {
variant request {
tiny(u8),
wide(u64),
}
roundtrip: func(request: request) -> request;
}
world primitive-variant-world {
export primitive-variant;
}
"#,
);
assert!(!world_uses_miden_core_types(&resolve, world));
}
#[test]
fn test_world_uses_miden_core_types_detects_imported_payload() {
let (resolve, world) = parse_test_world(
r#"
package miden:core-type-variant@0.1.0;
use miden:base/core-types@1.0.0;
interface core-type-variant {
use core-types.{word};
variant request {
elements(word),
amount(u64),
}
roundtrip: func(request: request) -> request;
}
world core-type-variant-world {
export core-type-variant;
}
"#,
);
assert!(world_uses_miden_core_types(&resolve, world));
}
}