1use proc_macro::TokenStream;
2use quote::quote;
3use syn::{
4 Attribute, Expr, Pat, Result as SynResult, Token, Type, TypePath,
5 parse::{Parse, ParseStream},
6 parse_macro_input, parse_str,
7};
8
9struct MatchWrapInput {
10 trait_type: Type,
11 expr: Expr,
12 arms: Vec<(Vec<Attribute>, Pat, Expr)>,
13}
14
15impl Parse for MatchWrapInput {
16 fn parse(input: ParseStream) -> SynResult<Self> {
17 let trait_type = input.parse()?;
18 input.parse::<Token![;]>()?;
19
20 let expr = input.parse()?;
21 input.parse::<Token![;]>()?;
22
23 let mut arms = Vec::new();
24 while !input.is_empty() {
25 let attrs = input.call(syn::Attribute::parse_outer)?;
26 let pat = Pat::parse_single(input)?;
27 input.parse::<Token![=>]>()?;
28 let arm_expr = input.parse()?;
29 arms.push((attrs, pat, arm_expr));
30
31 if input.peek(Token![,]) {
32 input.parse::<Token![,]>()?;
33 }
34 }
35
36 Ok(MatchWrapInput {
37 trait_type,
38 expr,
39 arms,
40 })
41 }
42}
43
44fn match_wrap(input: TokenStream, container_type: &str) -> TokenStream {
45 const DIVERGE_ATTR: &str = "diverges";
46 let input = parse_macro_input!(input as MatchWrapInput);
47 let trait_type = &input.trait_type;
48 let container: TypePath = parse_str(container_type).expect("");
49
50 let expr = &input.expr;
51
52 let arms = input.arms.iter().map(|(attrs, pat, arm_expr)| {
53 let is_diverging = attrs.iter().any(|attr| attr.path().is_ident(DIVERGE_ATTR));
54
55 if is_diverging {
56 quote! {#pat => #arm_expr}
57 } else {
58 quote! { #pat => #container::new(#arm_expr) as #container<#trait_type> }
59 }
60 });
61
62 quote! { match #expr { #(#arms,)* } }.into()
63}
64
65#[proc_macro]
66pub fn match_box(input: TokenStream) -> TokenStream {
67 match_wrap(input, "::std::boxed::Box")
68}
69
70#[proc_macro]
71pub fn match_arc(input: TokenStream) -> TokenStream {
72 match_wrap(input, "::std::sync::Arc")
73}
74
75#[proc_macro]
76pub fn match_rc(input: TokenStream) -> TokenStream {
77 match_wrap(input, "::std::rc::Rc")
78}