Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
219 changes: 219 additions & 0 deletions crates/core/src/chainable_method.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,219 @@
use anyhow::{Result, bail};
use std::collections::HashSet;
use std::fmt;
use wit_parser::{Function, FunctionKind, Resolve, WorldKey};

/// Structure used to parse the command line argument `--chainable-method` consistently
/// across guest generators.
#[cfg_attr(feature = "clap", derive(clap::Parser))]
#[cfg_attr(feature = "serde", derive(serde::Deserialize))]
#[derive(Clone, Default, Debug)]
pub struct ChainableMethodFilterSet {
/// Determines which resource methods should have chaining enabled.
/// Chaining takes a WIT method import returning nothing, and modifies bindgen
/// in a language-dependent way to return `self` in the glue code. This does
/// not affect the ABI in any way.
///
/// This option can be passed multiple times and additionally accepts
/// comma-separated values for each option passed. Each individual argument
/// passed here can be one of:
///
/// - `all` - all applicable methods will be chainable
/// - `foo:bar/baz#my-resource` - enable chaining for all methods in a resource
/// - `foo:bar/baz#my-resource.some-method` - enable chaining for particular method
///
/// Each filter may also have one of two modifier prefixes:
/// - `-` - inverts the selection; e.g. `-all` will disable chaining for all
/// - `&` - makes the chainable return `&Self` instead of `Self` (borrowing)
///
/// For instance, `&foo:bar/baz#my-resource` will make all methods in said resource
/// borrowing chainable, while `-foo:bar/baz#my-resource.some-method` will disable it
/// for that particular method.
///
/// Options are processed in the order they are passed here, so if a method
/// matches two directives passed the least-specific one should be last.
#[cfg_attr(
feature = "clap",
arg(
long = "chainable-methods",
value_parser = parse_chainable_method,
value_delimiter =',',
value_name = "FILTER",
),
)]
chainable_methods: Vec<ChainableMethod>,

#[cfg_attr(feature = "clap", arg(skip))]
#[cfg_attr(feature = "serde", serde(skip))]
used_options: HashSet<usize>,
}

#[cfg(feature = "clap")]
fn parse_chainable_method(s: &str) -> Result<ChainableMethod, String> {
Ok(ChainableMethod::parse(s))
}

#[derive(Clone, Copy, Debug)]
pub enum ChainingMode {
Owning,
Borrowing,
}

impl ChainableMethodFilterSet {
/// Returns a set where all functions should be chainable or not depending on
/// `enable` provided.
pub fn all(mode: ChainingMode) -> ChainableMethodFilterSet {
ChainableMethodFilterSet {
chainable_methods: vec![ChainableMethod {
mode: Some(mode),
filter: ChainableMethodFilter::All,
}],
used_options: HashSet::new(),
}
}

/// Returns whether the `func` provided should be made chainable
pub fn should_be_chainable(
&mut self,
resolve: &Resolve,
interface: Option<&WorldKey>,
func: &Function,
is_import: bool,
) -> Option<ChainingMode> {
if !is_import {
return None;
}

if func.result.is_some() {
return None;
}

match func.kind {
FunctionKind::AsyncMethod(resource) | FunctionKind::Method(resource) => {
let interface_name = match interface.map(|key| resolve.name_world_key(key)) {
Some(str) => str + "#",
None => "".into(),
};

let resource_name_to_test = format!(
"{}{}",
interface_name,
resolve.types[resource].name.as_ref().unwrap()
);

let method_name_to_test = format!("{}{}", interface_name, func.name);

for (i, opt) in self.chainable_methods.iter().enumerate() {
match &opt.filter {
ChainableMethodFilter::All => {
self.used_options.insert(i);
return opt.mode;
}
ChainableMethodFilter::Resource(s) => {
if *s == resource_name_to_test {
self.used_options.insert(i);
return opt.mode;
}
}
ChainableMethodFilter::Method(s) => {
if *s == method_name_to_test {
self.used_options.insert(i);
return opt.mode;
}
}
};
}

return None;
}
_ => {
return None;
}
}
}

/// Intended to be used in the header comment of generated code to help
/// indicate what options were specified.
pub fn debug_opts(&self) -> impl Iterator<Item = String> + '_ {
self.chainable_methods.iter().map(|opt| opt.to_string())
}

/// Tests whether all `--chainable-method` options were used throughout bindings
/// generation, returning an error if any were unused.
pub fn ensure_all_used(&self) -> Result<()> {
for (i, opt) in self.chainable_methods.iter().enumerate() {
if self.used_options.contains(&i) {
continue;
}
if !matches!(opt.filter, ChainableMethodFilter::All) {
bail!("unused chainable option: {opt}");
}
}
Ok(())
}

/// Pushes a new option into this set.
pub fn push(&mut self, directive: &str) {
self.chainable_methods
.push(ChainableMethod::parse(directive));
}
}

#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Deserialize))]
struct ChainableMethod {
mode: Option<ChainingMode>,
filter: ChainableMethodFilter,
}

impl ChainableMethod {
fn parse(s: &str) -> ChainableMethod {
let (s, mode) = match s.strip_prefix('-') {
Some(s) => (s, None),
None => match s.strip_prefix('&') {
Some(s) => (s, Some(ChainingMode::Borrowing)),
None => (s, Some(ChainingMode::Owning)),
},
};
let filter = match s {
"all" => ChainableMethodFilter::All,
other => {
if other.contains("[method]") {
ChainableMethodFilter::Method(other.to_string())
} else {
ChainableMethodFilter::Resource(other.to_string())
}
}
};
ChainableMethod { mode, filter }
}
}

impl fmt::Display for ChainableMethod {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.mode {
Some(ChainingMode::Owning) => {}
Some(ChainingMode::Borrowing) => write!(f, "&")?,
None => write!(f, "-")?,
};
self.filter.fmt(f)
}
}

#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Deserialize))]
enum ChainableMethodFilter {
All,
Resource(String),
Method(String),
}

impl fmt::Display for ChainableMethodFilter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ChainableMethodFilter::All => write!(f, "all"),
ChainableMethodFilter::Resource(s) => write!(f, "{s}"),
ChainableMethodFilter::Method(s) => write!(f, "{s}"),
}
}
}
2 changes: 2 additions & 0 deletions crates/core/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@ mod path;
pub use path::name_package_module;
mod async_;
pub use async_::AsyncFilterSet;
mod chainable_method;
pub use chainable_method::{ChainableMethodFilterSet, ChainingMode};

#[derive(Default, Copy, Clone, PartialEq, Eq, Debug)]
pub enum Direction {
Expand Down
31 changes: 23 additions & 8 deletions crates/guest-rust/macro/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,9 @@ use std::sync::atomic::{AtomicUsize, Ordering::Relaxed};
use syn::parse::{Error, Parse, ParseStream, Result};
use syn::punctuated::Punctuated;
use syn::{Token, braced, token};
use wit_bindgen_core::AsyncFilterSet;
use wit_bindgen_core::WorldGenerator;
use wit_bindgen_core::wit_parser::{PackageId, Resolve, WorldId};
use wit_bindgen_core::{AsyncFilterSet, ChainableMethodFilterSet};
use wit_bindgen_rust::{Opts, Ownership, WithOption};

#[proc_macro]
Expand Down Expand Up @@ -66,6 +66,7 @@ impl Parse for Config {
let mut source = None;
let mut features = Vec::new();
let mut async_configured = false;
let mut method_chaining_configured = false;
let mut debug = false;

if input.peek(token::Brace) {
Expand Down Expand Up @@ -169,8 +170,15 @@ impl Parse for Config {
async_configured = true;
opts.async_ = val;
}
Opt::EnableMethodChaining(enable) => {
opts.enable_method_chaining = enable.value();
Opt::ChainableMethods(val, span) => {
if method_chaining_configured {
return Err(Error::new(
span,
"cannot specify second method chaining config",
));
}
method_chaining_configured = true;
opts.chainable_methods = val;
}
Opt::MergeStructurallyEqualTypes(enable) => {
opts.merge_structurally_equal_types = Some(Some(enable.value()))
Expand Down Expand Up @@ -330,7 +338,7 @@ mod kw {
syn::custom_keyword!(disable_custom_section_link_helpers);
syn::custom_keyword!(imports);
syn::custom_keyword!(debug);
syn::custom_keyword!(enable_method_chaining);
syn::custom_keyword!(chainable_methods);
syn::custom_keyword!(merge_structurally_equal_types);
}

Expand Down Expand Up @@ -414,7 +422,7 @@ enum Opt {
DisableCustomSectionLinkHelpers(syn::LitBool),
Async(AsyncFilterSet, Span),
Debug(syn::LitBool),
EnableMethodChaining(syn::LitBool),
ChainableMethods(ChainableMethodFilterSet, Span),
MergeStructurallyEqualTypes(syn::LitBool),
}

Expand Down Expand Up @@ -600,10 +608,17 @@ impl Parse for Opt {
input.parse::<kw::debug>()?;
input.parse::<Token![:]>()?;
Ok(Opt::Debug(input.parse()?))
} else if l.peek(kw::enable_method_chaining) {
input.parse::<kw::enable_method_chaining>()?;
} else if l.peek(kw::chainable_methods) {
let span = input.parse::<kw::chainable_methods>()?.span;
input.parse::<Token![:]>()?;
Ok(Opt::EnableMethodChaining(input.parse()?))

let mut set = ChainableMethodFilterSet::default();
let contents;
syn::bracketed!(contents in input);
for val in contents.parse_terminated(|p| p.parse::<syn::LitStr>(), Token![,])? {
set.push(&val.value());
}
Ok(Opt::ChainableMethods(set, span))
} else if l.peek(Token![async]) {
let span = input.parse::<Token![async]>()?.span;
input.parse::<Token![:]>()?;
Expand Down
10 changes: 5 additions & 5 deletions crates/rust/src/bindgen.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ use heck::*;
use std::fmt::Write as _;
use std::mem;
use wit_bindgen_core::abi::{Bindgen, Instruction, LiftLower, WasmType};
use wit_bindgen_core::{Source, dealias, uwrite, uwriteln, wit_parser::*};
use wit_bindgen_core::{ChainingMode, Source, dealias, uwrite, uwriteln, wit_parser::*};

pub(super) struct FunctionBindgen<'a, 'b> {
pub r#gen: &'b mut InterfaceGenerator<'a>,
Expand All @@ -21,7 +21,7 @@ pub(super) struct FunctionBindgen<'a, 'b> {
pub import_return_pointer_area_align: Alignment,
pub handle_decls: Vec<String>,
always_owned: bool,
return_self: bool,
return_self: Option<ChainingMode>,
}

pub const POINTER_SIZE_EXPRESSION: &str = "::core::mem::size_of::<*const u8>()";
Expand All @@ -32,7 +32,7 @@ impl<'a, 'b> FunctionBindgen<'a, 'b> {
params: Vec<String>,
wasm_import_module: &'b str,
always_owned: bool,
return_self: bool,
return_self: Option<ChainingMode>,
) -> FunctionBindgen<'a, 'b> {
FunctionBindgen {
r#gen,
Expand Down Expand Up @@ -1054,11 +1054,11 @@ impl Bindgen for FunctionBindgen<'_, '_> {
}

Instruction::Return { amt, .. } => {
assert!(!self.return_self || *amt == 0);
assert!(self.return_self.is_none() || *amt == 0);

match amt {
0 => {
if self.return_self {
if self.return_self.is_some() {
self.push_str("self\n");
}
}
Expand Down
Loading
Loading