Skip to main content

delegate_attr/
lib.rs

1//! Attribute proc-macro to delegate method to a field.
2//!
3//! ## Examples
4//!
5//! ### Delegate `impl` block
6//!
7//! ```
8//! use delegate_attr::delegate;
9//!
10//! struct Foo(String);
11//!
12//! #[delegate(self.0)]
13//! impl Foo {
14//!     fn as_str(&self) -> &str {}
15//!     fn into_bytes(self) -> Vec<u8> {}
16//! }
17//!
18//! let foo = Foo("hello".to_owned());
19//! assert_eq!(foo.as_str(), "hello");
20//! assert_eq!(foo.into_bytes(), b"hello");
21//! ```
22//!
23//! ### Delegate trait `impl`
24//!
25//! ```
26//! # use delegate_attr::delegate;
27//!
28//! struct Iter(std::vec::IntoIter<u8>);
29//!
30//! #[delegate(self.0)]
31//! impl Iterator for Iter {
32//!     type Item = u8;
33//!     fn next(&mut self) -> Option<u8> {}
34//!     fn count(self) -> usize {}
35//!     fn size_hint(&self) -> (usize, Option<usize>) {}
36//!     fn last(self) -> Option<u8> {}
37//! }
38//!
39//! let iter = Iter(vec![1, 2, 4, 8].into_iter());
40//! assert_eq!(iter.count(), 4);
41//! let iter = Iter(vec![1, 2, 4, 8].into_iter());
42//! assert_eq!(iter.last(), Some(8));
43//! let iter = Iter(vec![1, 2, 4, 8].into_iter());
44//! assert_eq!(iter.sum::<u8>(), 15);
45//! ```
46//!
47//! ### With more complicated target
48//!
49//! ```
50//! # use delegate_attr::delegate;
51//! # use std::cell::RefCell;
52//! struct Foo<T> {
53//!     inner: RefCell<Vec<T>>,
54//! }
55//!
56//! #[delegate(self.inner.borrow())]
57//! impl<T> Foo<T> {
58//!     fn len(&self) -> usize {}
59//! }
60//!
61//! #[delegate(self.inner.borrow_mut())]
62//! impl<T> Foo<T> {
63//!     fn push(&self, value: T) {}
64//! }
65//!
66//! #[delegate(self.inner.into_inner())]
67//! impl<T> Foo<T> {
68//!     fn into_boxed_slice(self) -> Box<[T]> {}
69//! }
70//!
71//! let foo = Foo { inner: RefCell::new(vec![1]) };
72//! assert_eq!(foo.len(), 1);
73//! foo.push(2);
74//! assert_eq!(foo.len(), 2);
75//! assert_eq!(foo.into_boxed_slice().as_ref(), &[1, 2]);
76//! ```
77//!
78//! ### `into` and `call` attribute
79//!
80//! ```
81//! # use delegate_attr::delegate;
82//! struct Inner;
83//! impl Inner {
84//!     pub fn method(&self, num: u32) -> u32 { num }
85//! }
86//!
87//! struct Wrapper { inner: Inner }
88//!
89//! #[delegate(self.inner)]
90//! impl Wrapper {
91//!     // calls method, converts result to u64
92//!     #[into]
93//!     pub fn method(&self, num: u32) -> u64 {}
94//!
95//!     // calls method, returns ()
96//!     #[call(method)]
97//!     pub fn method_noreturn(&self, num: u32) {}
98//! }
99//! ```
100//!
101//! ### Delegate single method
102//!
103//! ```
104//! # use delegate_attr::delegate;
105//! struct Foo<T>(Vec<T>);
106//!
107//! impl<T> Foo<T> {
108//!     #[delegate(self.0)]
109//!     fn len(&self) -> usize {}
110//! }
111//!
112//! let foo = Foo(vec![1]);
113//! assert_eq!(foo.len(), 1);
114//! ```
115
116extern crate proc_macro;
117
118use proc_macro::TokenStream as RawTokenStream;
119use proc_macro2::{Group, Ident, TokenStream, TokenTree};
120use quote::{quote, quote_spanned, ToTokens};
121use syn::spanned::Spanned;
122use syn::{parse_macro_input, Expr, FnArg, ImplItem, ImplItemFn, ItemImpl, Meta, Pat, ReturnType};
123
124#[proc_macro_attribute]
125pub fn delegate(attr: RawTokenStream, item: RawTokenStream) -> RawTokenStream {
126    let receiver = parse_macro_input!(attr as Expr);
127    delegate_input(item.into(), &receiver).into()
128}
129
130fn delegate_input(input: TokenStream, receiver: &Expr) -> TokenStream {
131    if let Ok(input) = syn::parse2::<ItemImpl>(input.clone()) {
132        return delegate_impl_block(input, receiver);
133    }
134    if let Ok(input) = syn::parse2::<ImplItemFn>(input.clone()) {
135        return delegate_fn(input, receiver);
136    }
137    let mut tokens = input.into_iter();
138    let first_non_attr_token = 'outer: loop {
139        match tokens.next() {
140            None => break None,
141            Some(TokenTree::Punct(p)) if p.as_char() == '#' => {}
142            Some(token) => break Some(token),
143        }
144        loop {
145            match tokens.next() {
146                None => break 'outer None,
147                Some(TokenTree::Punct(_)) => {}
148                Some(TokenTree::Group(_)) => continue 'outer,
149                Some(token) => break 'outer Some(token),
150            }
151        }
152    };
153    if let Some(token) = first_non_attr_token {
154        let msg = match &token {
155            TokenTree::Ident(ident) if ident == "impl" => "invalid impl block for #[delegate]",
156            TokenTree::Ident(ident) if ident == "fn" => "invalid method for #[delegate]",
157            _ => "expected an impl block or method inside impl block",
158        };
159        quote_spanned! { token.span() => compile_error!(#msg); }
160    } else {
161        panic!("unexpected eof")
162    }
163}
164
165fn delegate_impl_block(input: ItemImpl, receiver: &Expr) -> TokenStream {
166    let ItemImpl {
167        attrs,
168        modifiers,
169        unsafety,
170        impl_token,
171        mut generics,
172        trait_,
173        self_ty,
174        brace_token: _,
175        items,
176    } = input;
177    let where_clause = generics.where_clause.take();
178    let trait_ = trait_.map(|(path, for_)| quote!(#path #for_));
179    let defaultness = &modifiers.defaultness;
180    let polarity = &modifiers.polarity;
181    let items = items.into_iter().map(|item| {
182        let func = match item {
183            ImplItem::Fn(f) => f,
184            _ => return item.into_token_stream(),
185        };
186        delegate_fn(func, receiver)
187    });
188
189    quote! {
190        #(#attrs)* #defaultness #unsafety #impl_token #generics #polarity #trait_ #self_ty #where_clause {
191            #(#items)*
192        }
193    }
194}
195
196fn delegate_fn(input: ImplItemFn, receiver: &Expr) -> TokenStream {
197    let ImplItemFn {
198        mut attrs,
199        vis,
200        modifiers,
201        sig,
202        block: _,
203    } = input;
204    let mut errors = TokenStream::new();
205    let defaultness = &modifiers.defaultness;
206    macro_rules! push_error {
207        ($error: expr) => {
208            errors.extend($error.into_compile_error())
209        };
210        ($span: expr, $msg: expr) => {
211            push_error!(syn::Error::new($span, $msg))
212        };
213    }
214    // Parse attributes.
215    let mut has_inline = false;
216    let mut has_into = false;
217    let mut call_name = None;
218    attrs.retain(|attr| {
219        let path = attr.path();
220        if path.is_ident("inline") {
221            has_inline = true;
222        } else if path.is_ident("into") {
223            match &attr.meta {
224                Meta::List(meta) => {
225                    push_error!(meta.delimiter.span().join(), "unexpected argument")
226                }
227                Meta::NameValue(meta) => push_error!(meta.eq_token.span, "unexpected argument"),
228                Meta::Path(_) => {}
229            }
230            if has_into {
231                push_error!(attr.span(), "duplicate #[into] attribute");
232            }
233            has_into = true;
234            return false;
235        } else if path.is_ident("call") {
236            match attr.parse_args::<Ident>() {
237                Ok(ident) => {
238                    if call_name.is_some() {
239                        push_error!(attr.span(), "duplicate #[call] attribute");
240                    }
241                    call_name = Some(ident);
242                }
243                Err(e) => push_error!(e),
244            }
245            return false;
246        }
247        true
248    });
249    // Mark method always inline if it's not otherwise specified.
250    let inline = if !has_inline {
251        quote!(#[inline(always)])
252    } else {
253        quote!()
254    };
255    let mut inputs = sig.inputs.iter();
256    // Extract the self token.
257    let self_token = match inputs.next() {
258        Some(FnArg::Receiver(receiver)) => receiver.self_token.to_token_stream(),
259        Some(FnArg::Typed(pat)) => match &*pat.pat {
260            Pat::Ident(ident) if ident.ident == "self" => ident.ident.to_token_stream(),
261            _ => {
262                push_error!(pat.span(), "expected self");
263                TokenStream::new()
264            }
265        },
266        None => {
267            push_error!(sig.paren_token.span.join(), "expected self");
268            TokenStream::new()
269        }
270    };
271    // List all parameters.
272    let args = inputs
273        .filter_map(|arg| match arg {
274            FnArg::Typed(pat) => match &*pat.pat {
275                Pat::Ident(ident) => Some(ident.to_token_stream()),
276                _ => {
277                    push_error!(pat.pat.span(), "expect an identifier");
278                    None
279                }
280            },
281            _ => {
282                push_error!(arg.span(), "unexpected argument");
283                None
284            }
285        })
286        .collect::<Vec<_>>();
287    // Return errors if any.
288    if !errors.is_empty() {
289        return errors;
290    } else {
291        // Drop it to ensure that we are not pushing anymore into it.
292        drop(errors);
293    }
294    // Generate method call.
295    let name = call_name.as_ref().unwrap_or(&sig.ident);
296    // Replace the self token in the receiver with the token we extract above to ensure it comes
297    // from the right hygiene context.
298    let receiver = replace_self(receiver.to_token_stream(), &self_token);
299    let body = quote! { #receiver.#name(#(#args),*) };
300    let body = match &sig.output {
301        ReturnType::Default => quote! { #body; },
302        ReturnType::Type(_, ty) if has_into => {
303            quote! { ::std::convert::Into::<#ty>::into(#body) }
304        }
305        _ => body,
306    };
307    quote! {
308        #(#attrs)* #inline #vis #defaultness #sig {
309            #body
310        }
311    }
312}
313
314fn replace_self(expr: TokenStream, self_token: &TokenStream) -> TokenStream {
315    expr.into_iter()
316        .map(|token| match token {
317            TokenTree::Ident(ident) if ident == "self" => self_token.clone(),
318            TokenTree::Group(group) => {
319                let delimiter = group.delimiter();
320                let stream = replace_self(group.stream(), self_token);
321                Group::new(delimiter, stream).into_token_stream()
322            }
323            _ => token.into_token_stream(),
324        })
325        .collect()
326}