Skip to main content

extend/
lib.rs

1//! Create extensions for types you don't own with [extension traits] but without the boilerplate.
2//!
3//! Example:
4//!
5//! ```rust
6//! use extend::ext;
7//!
8//! #[ext]
9//! impl<T: Ord> Vec<T> {
10//!     fn sorted(mut self) -> Self {
11//!         self.sort();
12//!         self
13//!     }
14//! }
15//!
16//! assert_eq!(
17//!     vec![1, 2, 3],
18//!     vec![2, 3, 1].sorted(),
19//! );
20//! ```
21//!
22//! # How does it work?
23//!
24//! Under the hood it generates a trait with methods in your `impl` and implements those for the
25//! type you specify. The code shown above expands roughly to:
26//!
27//! ```rust
28//! trait VecExt<T: Ord> {
29//!     fn sorted(self) -> Self;
30//! }
31//!
32//! impl<T: Ord> VecExt<T> for Vec<T> {
33//!     fn sorted(mut self) -> Self {
34//!         self.sort();
35//!         self
36//!     }
37//! }
38//! ```
39//!
40//! # Supported items
41//!
42//! Extensions can contain methods or associated constants:
43//!
44//! ```rust
45//! use extend::ext;
46//!
47//! #[ext]
48//! impl String {
49//!     const CONSTANT: &'static str = "FOO";
50//!
51//!     fn method() {
52//!         // ...
53//!         # todo!()
54//!     }
55//! }
56//! ```
57//!
58//! # Configuration
59//!
60//! You can configure:
61//!
62//! - The visibility of the trait. Use `pub impl ...` to generate `pub trait ...`. The default
63//! visibility is private.
64//! - The name of the generated extension trait. Example: `#[ext(name = MyExt)]`. By default we
65//! generate a name based on what you extend.
66//! - Which supertraits the generated extension trait should have. Default is no supertraits.
67//! Example: `#[ext(supertraits = Default + Clone)]`.
68//!
69//! More examples:
70//!
71//! ```rust
72//! use extend::ext;
73//!
74//! #[ext(name = SortedVecExt)]
75//! impl<T: Ord> Vec<T> {
76//!     fn sorted(mut self) -> Self {
77//!         self.sort();
78//!         self
79//!     }
80//! }
81//!
82//! #[ext]
83//! pub(crate) impl i32 {
84//!     fn double(self) -> i32 {
85//!         self * 2
86//!     }
87//! }
88//!
89//! #[ext(name = ResultSafeUnwrapExt)]
90//! pub impl<T> Result<T, std::convert::Infallible> {
91//!     fn safe_unwrap(self) -> T {
92//!         match self {
93//!             Ok(t) => t,
94//!             Err(_) => unreachable!(),
95//!         }
96//!     }
97//! }
98//!
99//! #[ext(supertraits = Default + Clone)]
100//! impl String {
101//!     fn my_length(self) -> usize {
102//!         self.len()
103//!     }
104//! }
105//! ```
106//!
107//! For backwards compatibility you can also declare the visibility as the first argument to `#[ext]`:
108//!
109//! ```
110//! use extend::ext;
111//!
112//! #[ext(pub)]
113//! impl i32 {
114//!     fn double(self) -> i32 {
115//!         self * 2
116//!     }
117//! }
118//! ```
119//!
120//! # async-trait compatibility
121//!
122//! Async extensions are supported via [async-trait](https://crates.io/crates/async-trait).
123//!
124//! Be aware that you need to add `#[async_trait]` _below_ `#[ext]`. Otherwise the `ext` macro
125//! cannot see the `#[async_trait]` attribute and pass it along in the generated code.
126//!
127//! Example:
128//!
129//! ```
130//! use extend::ext;
131//! use async_trait::async_trait;
132//!
133//! #[ext]
134//! #[async_trait]
135//! impl String {
136//!     async fn read_file() -> String {
137//!         // ...
138//!         # todo!()
139//!     }
140//! }
141//! ```
142//!
143//! # Other attributes
144//!
145//! Other attributes provided _below_ `#[ext]` will be passed along to both the generated trait and
146//! the implementation. See [async-trait compatibility](#async-trait-compatibility) above for an
147//! example.
148//!
149//! [extension traits]: https://dev.to/matsimitsu/extending-existing-functionality-in-rust-with-traits-in-rust-3622
150
151#![allow(clippy::let_and_return)]
152#![deny(unused_variables, dead_code, unused_must_use, unused_imports)]
153
154use proc_macro2::TokenStream;
155use quote::{format_ident, quote, ToTokens};
156use std::convert::{TryFrom, TryInto};
157use syn::{
158    parse::{self, Parse, ParseStream},
159    parse_macro_input, parse_quote,
160    punctuated::Punctuated,
161    spanned::Spanned,
162    token::{Plus, Semi},
163    Ident, ImplItem, ItemImpl, Result, Token, TraitItemConst, TraitItemFn, Type, TypeArray,
164    TypeBareFn, TypeGroup, TypeNever, TypeParamBound, TypeParen, TypePath, TypePtr, TypeReference,
165    TypeSlice, TypeTraitObject, TypeTuple, Visibility,
166};
167
168#[derive(Debug)]
169struct Input {
170    item_impl: ItemImpl,
171    vis: Option<Visibility>,
172}
173
174impl Parse for Input {
175    fn parse(input: ParseStream) -> syn::Result<Self> {
176        let mut attributes = Vec::new();
177        if input.peek(syn::Token![#]) {
178            attributes.extend(syn::Attribute::parse_outer(input)?);
179        }
180
181        let vis = input
182            .parse::<Visibility>()
183            .ok()
184            .filter(|vis| vis != &Visibility::Inherited);
185
186        let mut item_impl = input.parse::<ItemImpl>()?;
187        item_impl.attrs.extend(attributes);
188
189        Ok(Self { item_impl, vis })
190    }
191}
192
193/// See crate docs for more info.
194#[proc_macro_attribute]
195#[allow(clippy::unneeded_field_pattern)]
196pub fn ext(
197    attr: proc_macro::TokenStream,
198    item: proc_macro::TokenStream,
199) -> proc_macro::TokenStream {
200    let item = parse_macro_input!(item as Input);
201    let config = parse_macro_input!(attr as Config);
202    match go(item, config) {
203        Ok(tokens) => tokens,
204        Err(err) => err.into_compile_error().into(),
205    }
206}
207
208/// Like [`ext`](macro@crate::ext) but always add `Sized` as a supertrait.
209///
210/// This is provided as a convenience for generating extension traits that require `Self: Sized`
211/// such as:
212///
213/// ```
214/// use extend::ext_sized;
215///
216/// #[ext_sized]
217/// impl i32 {
218///     fn requires_sized(self) -> Option<Self> {
219///         Some(self)
220///     }
221/// }
222/// ```
223#[proc_macro_attribute]
224#[allow(clippy::unneeded_field_pattern)]
225pub fn ext_sized(
226    attr: proc_macro::TokenStream,
227    item: proc_macro::TokenStream,
228) -> proc_macro::TokenStream {
229    let item = parse_macro_input!(item as Input);
230    let mut config: Config = parse_macro_input!(attr as Config);
231
232    config.supertraits = if let Some(supertraits) = config.supertraits.take() {
233        Some(parse_quote!(#supertraits + Sized))
234    } else {
235        Some(parse_quote!(Sized))
236    };
237
238    match go(item, config) {
239        Ok(tokens) => tokens,
240        Err(err) => err.into_compile_error().into(),
241    }
242}
243
244fn go(item: Input, mut config: Config) -> Result<proc_macro::TokenStream> {
245    if let Some(vis) = item.vis {
246        if config.visibility != Visibility::Inherited {
247            return Err(syn::Error::new(
248                config.visibility.span(),
249                "Cannot set visibility on `#[ext]` and `impl` block",
250            ));
251        }
252
253        config.visibility = vis;
254    }
255
256    let ItemImpl {
257        attrs,
258        unsafety,
259        generics,
260        trait_,
261        self_ty,
262        items,
263        // What is defaultness?
264        defaultness: _,
265        impl_token: _,
266        brace_token: _,
267    } = item.item_impl;
268
269    if let Some((_, path, _)) = trait_ {
270        return Err(syn::Error::new(
271            path.span(),
272            "Trait impls cannot be used for #[ext]",
273        ));
274    }
275
276    let self_ty = parse_self_ty(&self_ty)?;
277
278    let ext_trait_name = if let Some(ext_trait_name) = config.ext_trait_name {
279        ext_trait_name
280    } else {
281        ext_trait_name(&self_ty)?
282    };
283
284    let MethodsAndConsts {
285        trait_methods,
286        trait_consts,
287    } = extract_allowed_items(&items)?;
288
289    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
290
291    let visibility = &config.visibility;
292
293    let mut all_supertraits = Vec::<TypeParamBound>::new();
294
295    if let Some(supertraits_from_config) = config.supertraits {
296        all_supertraits.extend(supertraits_from_config);
297    }
298
299    let supertraits_quoted = if all_supertraits.is_empty() {
300        quote! {}
301    } else {
302        let supertraits_quoted = punctuated_from_iter::<_, _, Plus>(all_supertraits);
303        quote! { : #supertraits_quoted }
304    };
305
306    let code = (quote! {
307        #[allow(non_camel_case_types)]
308        #(#attrs)*
309        #visibility
310        #unsafety
311        trait #ext_trait_name #impl_generics #supertraits_quoted #where_clause {
312            #(
313                #trait_consts
314            )*
315
316            #(
317                #[allow(
318                    patterns_in_fns_without_body,
319                    clippy::inline_fn_without_body,
320                    unused_attributes
321                )]
322                #trait_methods
323            )*
324        }
325
326        #(#attrs)*
327        impl #impl_generics #ext_trait_name #ty_generics for #self_ty #where_clause {
328            #(#items)*
329        }
330    })
331    .into();
332
333    Ok(code)
334}
335
336#[derive(Debug, Clone)]
337enum ExtType<'a> {
338    Array(&'a TypeArray),
339    Group(&'a TypeGroup),
340    Never(&'a TypeNever),
341    Paren(&'a TypeParen),
342    Path(&'a TypePath),
343    Ptr(&'a TypePtr),
344    Reference(&'a TypeReference),
345    Slice(&'a TypeSlice),
346    Tuple(&'a TypeTuple),
347    BareFn(&'a TypeBareFn),
348    TraitObject(&'a TypeTraitObject),
349}
350
351#[allow(clippy::wildcard_in_or_patterns)]
352fn parse_self_ty(self_ty: &Type) -> Result<ExtType> {
353    let ty = match self_ty {
354        Type::Array(inner) => ExtType::Array(inner),
355        Type::Group(inner) => ExtType::Group(inner),
356        Type::Never(inner) => ExtType::Never(inner),
357        Type::Paren(inner) => ExtType::Paren(inner),
358        Type::Path(inner) => ExtType::Path(inner),
359        Type::Ptr(inner) => ExtType::Ptr(inner),
360        Type::Reference(inner) => ExtType::Reference(inner),
361        Type::Slice(inner) => ExtType::Slice(inner),
362        Type::Tuple(inner) => ExtType::Tuple(inner),
363        Type::BareFn(inner) => ExtType::BareFn(inner),
364        Type::TraitObject(inner) => ExtType::TraitObject(inner),
365
366        Type::ImplTrait(_) | Type::Infer(_) | Type::Macro(_) | Type::Verbatim(_) | _ => {
367            return Err(syn::Error::new(
368                self_ty.span(),
369                "#[ext] is not supported for this kind of type",
370            ))
371        }
372    };
373    Ok(ty)
374}
375
376impl<'a> TryFrom<&'a Type> for ExtType<'a> {
377    type Error = syn::Error;
378
379    fn try_from(inner: &'a Type) -> Result<ExtType<'a>> {
380        parse_self_ty(inner)
381    }
382}
383
384impl<'a> ToTokens for ExtType<'a> {
385    fn to_tokens(&self, tokens: &mut TokenStream) {
386        match self {
387            ExtType::Array(inner) => inner.to_tokens(tokens),
388            ExtType::Group(inner) => inner.to_tokens(tokens),
389            ExtType::Never(inner) => inner.to_tokens(tokens),
390            ExtType::Paren(inner) => inner.to_tokens(tokens),
391            ExtType::Path(inner) => inner.to_tokens(tokens),
392            ExtType::Ptr(inner) => inner.to_tokens(tokens),
393            ExtType::Reference(inner) => inner.to_tokens(tokens),
394            ExtType::Slice(inner) => inner.to_tokens(tokens),
395            ExtType::Tuple(inner) => inner.to_tokens(tokens),
396            ExtType::BareFn(inner) => inner.to_tokens(tokens),
397            ExtType::TraitObject(inner) => inner.to_tokens(tokens),
398        }
399    }
400}
401
402fn ext_trait_name(self_ty: &ExtType) -> Result<Ident> {
403    fn inner_self_ty(self_ty: &ExtType) -> Result<Ident> {
404        match self_ty {
405            ExtType::Path(inner) => find_and_combine_idents(inner),
406            ExtType::Reference(inner) => {
407                let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
408                if inner.mutability.is_some() {
409                    Ok(format_ident!("RefMut{}", name))
410                } else {
411                    Ok(format_ident!("Ref{}", name))
412                }
413            }
414            ExtType::Array(inner) => {
415                let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
416                Ok(format_ident!("ListOf{}", name))
417            }
418            ExtType::Group(inner) => {
419                let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
420                Ok(format_ident!("Group{}", name))
421            }
422            ExtType::Paren(inner) => {
423                let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
424                Ok(format_ident!("Paren{}", name))
425            }
426            ExtType::Ptr(inner) => {
427                let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
428                Ok(format_ident!("PointerTo{}", name))
429            }
430            ExtType::Slice(inner) => {
431                let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
432                Ok(format_ident!("SliceOf{}", name))
433            }
434            ExtType::Tuple(inner) => {
435                let mut name = format_ident!("TupleOf");
436                for elem in &inner.elems {
437                    name = format_ident!("{}{}", name, inner_self_ty(&elem.try_into()?)?);
438                }
439                Ok(name)
440            }
441            ExtType::Never(_) => Ok(format_ident!("Never")),
442            ExtType::BareFn(inner) => {
443                let mut name = format_ident!("BareFn");
444                for input in inner.inputs.iter() {
445                    name = format_ident!("{}{}", name, inner_self_ty(&(&input.ty).try_into()?)?);
446                }
447                match &inner.output {
448                    syn::ReturnType::Default => {
449                        name = format_ident!("{}Unit", name);
450                    }
451                    syn::ReturnType::Type(_, ty) => {
452                        name = format_ident!("{}{}", name, inner_self_ty(&(&**ty).try_into()?)?);
453                    }
454                }
455                Ok(name)
456            }
457            ExtType::TraitObject(inner) => {
458                let mut name = format_ident!("TraitObject");
459                for bound in inner.bounds.iter() {
460                    match bound {
461                        TypeParamBound::Trait(bound) => {
462                            for segment in bound.path.segments.iter() {
463                                name = format_ident!("{}{}", name, segment.ident);
464                            }
465                        }
466                        TypeParamBound::Lifetime(lifetime) => {
467                            name = format_ident!("{}{}", name, lifetime.ident);
468                        }
469                        other => {
470                            return Err(syn::Error::new(other.span(), "unsupported bound"));
471                        }
472                    }
473                }
474                Ok(name)
475            }
476        }
477    }
478
479    Ok(format_ident!("{}Ext", inner_self_ty(self_ty)?))
480}
481
482fn find_and_combine_idents(type_path: &TypePath) -> Result<Ident> {
483    use syn::visit::{self, Visit};
484
485    struct IdentVisitor<'a>(Vec<&'a Ident>);
486
487    impl<'a> Visit<'a> for IdentVisitor<'a> {
488        fn visit_ident(&mut self, i: &'a Ident) {
489            self.0.push(i);
490        }
491    }
492
493    let mut visitor = IdentVisitor(Vec::new());
494    visit::visit_type_path(&mut visitor, type_path);
495    let idents = visitor.0;
496
497    if idents.is_empty() {
498        Err(syn::Error::new(type_path.span(), "Empty type path"))
499    } else {
500        let start = &idents[0].span();
501        let combined_span = idents
502            .iter()
503            .map(|i| i.span())
504            .fold(*start, |a, b| a.join(b).unwrap_or(a));
505
506        let combined_name = idents.iter().map(|i| i.to_string()).collect::<String>();
507
508        Ok(Ident::new(&combined_name, combined_span))
509    }
510}
511
512#[derive(Debug, Default)]
513struct MethodsAndConsts {
514    trait_methods: Vec<TraitItemFn>,
515    trait_consts: Vec<TraitItemConst>,
516}
517
518#[allow(clippy::wildcard_in_or_patterns)]
519fn extract_allowed_items(items: &[ImplItem]) -> Result<MethodsAndConsts> {
520    let mut acc = MethodsAndConsts::default();
521    for item in items {
522        match item {
523            ImplItem::Fn(method) => acc.trait_methods.push(TraitItemFn {
524                attrs: method.attrs.clone(),
525                sig: {
526                    let mut sig = method.sig.clone();
527                    sig.inputs = sig
528                        .inputs
529                        .into_iter()
530                        .map(|fn_arg| match fn_arg {
531                            syn::FnArg::Receiver(recv) => syn::FnArg::Receiver(recv),
532                            syn::FnArg::Typed(mut pat_type) => {
533                                pat_type.pat = Box::new(match *pat_type.pat {
534                                    syn::Pat::Ident(pat_ident) => syn::Pat::Ident(pat_ident),
535                                    _ => {
536                                        parse_quote!(_)
537                                    }
538                                });
539                                syn::FnArg::Typed(pat_type)
540                            }
541                        })
542                        .collect();
543                    sig
544                },
545                default: None,
546                semi_token: Some(Semi::default()),
547            }),
548            ImplItem::Const(const_) => acc.trait_consts.push(TraitItemConst {
549                attrs: const_.attrs.clone(),
550                generics: const_.generics.clone(),
551                const_token: Default::default(),
552                ident: const_.ident.clone(),
553                colon_token: Default::default(),
554                ty: const_.ty.clone(),
555                default: None,
556                semi_token: Default::default(),
557            }),
558            ImplItem::Type(_) => {
559                return Err(syn::Error::new(
560                    item.span(),
561                    "Associated types are not allowed in #[ext] impls",
562                ))
563            }
564            ImplItem::Macro(_) => {
565                return Err(syn::Error::new(
566                    item.span(),
567                    "Macros are not allowed in #[ext] impls",
568                ))
569            }
570            ImplItem::Verbatim(_) | _ => {
571                return Err(syn::Error::new(item.span(), "Not allowed in #[ext] impls"))
572            }
573        }
574    }
575    Ok(acc)
576}
577
578#[derive(Debug)]
579struct Config {
580    ext_trait_name: Option<Ident>,
581    visibility: Visibility,
582    supertraits: Option<Punctuated<TypeParamBound, Plus>>,
583}
584
585impl Parse for Config {
586    fn parse(input: ParseStream) -> parse::Result<Self> {
587        let mut config = Config::default();
588
589        if let Ok(visibility) = input.parse::<Visibility>() {
590            config.visibility = visibility;
591        }
592
593        input.parse::<Token![,]>().ok();
594
595        while !input.is_empty() {
596            let ident = input.parse::<Ident>()?;
597            input.parse::<Token![=]>()?;
598
599            match &*ident.to_string() {
600                "name" => {
601                    config.ext_trait_name = Some(input.parse()?);
602                }
603                "supertraits" => {
604                    config.supertraits =
605                        Some(Punctuated::<TypeParamBound, Plus>::parse_terminated(input)?);
606                }
607                _ => return Err(syn::Error::new(ident.span(), "Unknown configuration name")),
608            }
609
610            input.parse::<Token![,]>().ok();
611        }
612
613        Ok(config)
614    }
615}
616
617impl Default for Config {
618    fn default() -> Self {
619        Self {
620            ext_trait_name: None,
621            visibility: Visibility::Inherited,
622            supertraits: None,
623        }
624    }
625}
626
627fn punctuated_from_iter<I, T, P>(i: I) -> Punctuated<T, P>
628where
629    P: Default,
630    I: IntoIterator<Item = T>,
631{
632    let mut iter = i.into_iter().peekable();
633    let mut acc = Punctuated::default();
634
635    while let Some(item) = iter.next() {
636        acc.push_value(item);
637
638        if iter.peek().is_some() {
639            acc.push_punct(P::default());
640        }
641    }
642
643    acc
644}
645
646#[cfg(test)]
647mod test {
648    #[allow(unused_imports)]
649    use super::*;
650
651    #[test]
652    fn test_ui() {
653        let t = trybuild::TestCases::new();
654        t.pass("tests/compile_pass/*.rs");
655        t.compile_fail("tests/compile_fail/*.rs");
656    }
657}