use proc_macro::TokenStream;
use quote::quote;
use syn::{
Expr, ItemFn, ItemImpl, ItemStruct, ReturnType, Stmt, parse_macro_input, visit_mut::VisitMut,
};
struct ReturnVisitor {
errors: Vec<syn::Error>,
}
impl ReturnVisitor {
fn new() -> Self {
Self { errors: Vec::new() }
}
}
impl VisitMut for ReturnVisitor {
fn visit_expr_mut(&mut self, expr: &mut Expr) {
match expr {
Expr::Return(return_expr) => {
if return_expr.expr.is_some() {
self.errors.push(syn::Error::new_spanned(
return_expr.expr.clone(),
"chainable functions should not have explicit return values",
));
}
return_expr.expr = Some(Box::new(syn::parse_quote!(self)));
}
_ => {
syn::visit_mut::visit_expr_mut(self, expr);
}
}
}
}
fn make_function_chainable(_attr: TokenStream, item: TokenStream) -> TokenStream {
let mut input_fn = parse_macro_input!(item as ItemFn);
if input_fn.sig.inputs.is_empty() {
let error = syn::Error::new_spanned(
&input_fn.sig.ident,
"chainable functions must have at least one argument",
)
.to_compile_error();
return quote! {
#error
#input_fn
}
.into();
}
let (is_ref, is_mut_self) = input_fn
.sig
.inputs
.first()
.map(|arg| match arg {
syn::FnArg::Receiver(recv) => (recv.reference.is_some(), recv.mutability.is_some()),
_ => (false, false),
})
.unwrap();
if !is_mut_self {
let error = syn::Error::new_spanned(
&input_fn.sig.inputs[0],
"expected '&mut self' or 'mut self'",
)
.to_compile_error();
return quote! {
#error
#input_fn
}
.into();
}
match &input_fn.sig.output {
ReturnType::Type(_, _) => {
let error = syn::Error::new_spanned(
&input_fn.sig.output,
"chainable functions should not have explicit return types",
)
.to_compile_error();
return quote! {
#error
#input_fn
}
.into();
}
ReturnType::Default => {} }
input_fn.sig.output = if is_ref {
syn::parse_quote!(-> &mut Self)
} else {
syn::parse_quote!(-> Self)
};
let mut visitor = ReturnVisitor::new();
visitor.visit_block_mut(&mut input_fn.block);
if !visitor.errors.is_empty() {
let error = visitor.errors[0].to_compile_error();
return quote! {
#error
#input_fn
}
.into();
}
if let Some(Stmt::Expr(Expr::Return(_), _)) = input_fn.block.stmts.last() {
} else {
let return_stmt: Stmt = syn::parse_quote!(return self;);
input_fn.block.stmts.push(return_stmt);
}
let result = quote! {
#input_fn
};
result.into()
}
fn make_struct_chainable(_attr: TokenStream, item: TokenStream) -> TokenStream {
let input_struct = parse_macro_input!(item as ItemStruct);
let methods: Vec<_> = input_struct
.fields
.iter()
.map(|field| {
let function_name = syn::Ident::new(
&format!("with_{}", field.ident.as_ref().unwrap()),
field.ident.as_ref().unwrap().span(),
);
let field_name = &field.ident;
let field_type = &field.ty;
quote! {
pub fn #function_name(mut self, value: #field_type) -> Self {
self.#field_name = value;
self
}
}
})
.collect();
let name = &input_struct.ident;
let result = quote! {
#input_struct
impl #name {
#(#methods)*
}
};
eprintln!("Generated methods: {}", result);
result.into()
}
#[proc_macro_attribute]
pub fn chainable(_attr: TokenStream, item: TokenStream) -> TokenStream {
match syn::parse::<ItemFn>(item.clone()) {
Ok(_) => return make_function_chainable(_attr, item),
Err(_) => {}
}
match syn::parse::<ItemStruct>(item.clone()) {
Ok(_) => return make_struct_chainable(_attr, item),
Err(_) => {}
}
syn::Error::new_spanned(
proc_macro2::TokenStream::from(item),
"chainable attribute can only be applied to functions or implementations",
)
.to_compile_error()
.into()
}