Skip to main content

trillium_macros/
handler.rs

1use proc_macro::TokenStream;
2use quote::quote;
3use std::{collections::HashSet, iter::once};
4use syn::{
5    Attribute, Data, DeriveInput, Error, Expr, ExprArray, ExprAssign, ExprPath, Field, Ident,
6    Index, Member, Meta, Path, Type, TypePath, WhereClause,
7    parse::{Parse, ParseStream},
8    parse_macro_input, parse_quote,
9    punctuated::Punctuated,
10    spanned::Spanned,
11    token::{Comma, Where},
12    visit::{Visit, visit_type_path},
13};
14
15fn is_required_generic_for_type(ty: &Type, generic: &Ident) -> bool {
16    struct PathVisitor<'g> {
17        generic: &'g Ident,
18        generic_is_required: bool,
19    }
20    impl<'g, 'ast> Visit<'ast> for PathVisitor<'g> {
21        fn visit_type_path(&mut self, node: &'ast TypePath) {
22            if node.qself.is_none()
23                && let Some(first_segment) = node.path.segments.first()
24                && first_segment.ident == *self.generic
25            {
26                self.generic_is_required = true;
27            }
28            visit_type_path(self, node);
29        }
30    }
31
32    let mut path_visitor = PathVisitor {
33        generic,
34        generic_is_required: false,
35    };
36
37    path_visitor.visit_type(ty);
38
39    path_visitor.generic_is_required
40}
41
42#[derive(Clone, Copy, PartialEq, Eq, Debug)]
43enum Override {
44    Run,
45    Init,
46    BeforeSend,
47    HasUpgrade,
48    Upgrade,
49    Name,
50}
51
52impl TryFrom<&Path> for Override {
53    type Error = Error;
54
55    fn try_from(path: &Path) -> Result<Self, Self::Error> {
56        if path.is_ident("run") {
57            Ok(Self::Run)
58        } else if path.is_ident("init") {
59            Ok(Self::Init)
60        } else if path.is_ident("before_send") {
61            Ok(Self::BeforeSend)
62        } else if path.is_ident("has_upgrade") {
63            Ok(Self::HasUpgrade)
64        } else if path.is_ident("upgrade") {
65            Ok(Self::Upgrade)
66        } else if path.is_ident("name") {
67            Ok(Self::Name)
68        } else {
69            Err(Error::new(
70                path.span(),
71                "unrecognized trillium::Handler method name",
72            ))
73        }
74    }
75}
76
77struct DeriveOptions {
78    overrides: Vec<Override>,
79    input: DeriveInput,
80    field: Field,
81    field_index: usize,
82}
83
84fn overrides<'a, I: Iterator<Item = &'a Expr>>(iter: I) -> syn::Result<Vec<Override>> {
85    iter.map(|expr| match expr {
86        Expr::Path(ExprPath { path, .. }) => path.try_into(),
87        _ => Err(Error::new(
88            expr.span(),
89            "unrecognized override. valid options are run, init, before_send, name, has_upgrade, \
90             and upgrade",
91        )),
92    })
93    .collect()
94}
95
96fn parse_attribute(attr: &Attribute) -> syn::Result<Option<Vec<Override>>> {
97    if attr.path().is_ident("handler") {
98        match &attr.meta {
99            Meta::Path(_) => Ok(Some(vec![])),
100            Meta::List(metalist) => {
101                let tokens = metalist.tokens.clone();
102                let ExprAssign { left, right, .. } = syn::parse(tokens.into())?;
103                match (*left, *right) {
104                    (Expr::Path(ExprPath { path: left, .. }), right @ Expr::Path(_))
105                        if left.is_ident("except") =>
106                    {
107                        Ok(Some(overrides(once(&right))?))
108                    }
109
110                    (
111                        Expr::Path(ExprPath { path: left, .. }),
112                        Expr::Array(ExprArray { elems: right, .. }),
113                    ) if left.is_ident("except") => Ok(Some(overrides(right.iter())?)),
114
115                    (_x, _y) => Err(Error::new(
116                        metalist.span(),
117                        "unrecognized #[handler] attributes",
118                    )),
119                }
120            }
121            Meta::NameValue(nv) => Err(Error::new(nv.span(), "unrecognized #[handler] attributes")),
122        }
123    } else {
124        Ok(None)
125    }
126}
127
128fn generics(field: &Field, input: &DeriveInput) -> Vec<Ident> {
129    input
130        .generics
131        .type_params()
132        .filter_map(|g| {
133            if is_required_generic_for_type(&field.ty, &g.ident) {
134                Some(g.ident.clone())
135            } else {
136                None
137            }
138        })
139        .collect::<HashSet<_>>()
140        .into_iter()
141        .collect()
142}
143
144impl Parse for DeriveOptions {
145    fn parse(input: ParseStream) -> syn::Result<Self> {
146        let input = DeriveInput::parse(input)?;
147        let Data::Struct(ds) = &input.data else {
148            return Err(Error::new(input.span(), "second error"));
149        };
150
151        for (field_index, field) in ds.fields.iter().enumerate() {
152            for attr in &field.attrs {
153                if let Some(overrides) = parse_attribute(attr)? {
154                    let field = field.clone();
155                    return Ok(Self {
156                        overrides,
157                        input,
158                        field,
159                        field_index,
160                    });
161                }
162            }
163        }
164
165        if ds.fields.len() == 1 {
166            let field = ds
167                .fields
168                .iter()
169                .next()
170                .expect("len == 1 should have one element")
171                .clone();
172            Ok(Self {
173                overrides: vec![],
174                input,
175                field,
176                field_index: 0,
177            })
178        } else {
179            Err(Error::new(
180                input.span(),
181                "Structs with more than one field need a #[handler] annotation",
182            ))
183        }
184    }
185}
186
187pub fn derive_handler(input: TokenStream) -> TokenStream {
188    let DeriveOptions {
189        overrides,
190        field,
191        input,
192        field_index,
193    } = parse_macro_input!(input as DeriveOptions);
194
195    let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
196
197    let generics = generics(&field, &input);
198
199    let struct_name = input.ident;
200
201    let mut where_clause = where_clause.map_or_else(
202        || WhereClause {
203            where_token: Where::default(),
204            predicates: Punctuated::new(),
205        },
206        |where_clause| where_clause.to_owned(),
207    );
208
209    for generic in generics {
210        where_clause
211            .predicates
212            .push_value(parse_quote! { #generic: trillium::Handler });
213        where_clause.predicates.push_punct(Comma::default());
214    }
215
216    where_clause
217        .predicates
218        .push_value(parse_quote! { Self: Send + Sync + 'static });
219
220    let handler = field
221        .ident
222        .map_or_else(|| Member::Unnamed(Index::from(field_index)), Member::Named);
223
224    let handler = quote!(self.#handler);
225
226    let run = if overrides.contains(&Override::Run) {
227        quote!(Self::run(&self, conn))
228    } else {
229        quote!(trillium::Handler::run(&#handler, conn))
230    };
231
232    let init = if overrides.contains(&Override::Init) {
233        quote!(Self::init(&mut self, info))
234    } else {
235        quote!(trillium::Handler::init(&mut #handler, info))
236    };
237
238    let before_send = if overrides.contains(&Override::BeforeSend) {
239        quote!(Self::before_send(&self, conn))
240    } else {
241        quote!(trillium::Handler::before_send(&#handler, conn))
242    };
243
244    let name = if overrides.contains(&Override::Name) {
245        quote!(Self::name(&self))
246    } else {
247        let name_string = struct_name.to_string();
248        quote!(format!("{} ({})", #name_string, trillium::Handler::name(&#handler)).into())
249    };
250
251    let has_upgrade = if overrides.contains(&Override::HasUpgrade) {
252        quote!(Self::has_upgrade(&self, upgrade))
253    } else {
254        quote!(trillium::Handler::has_upgrade(&#handler, upgrade))
255    };
256
257    let upgrade = if overrides.contains(&Override::Upgrade) {
258        quote!(Self::upgrade(&self, upgrade))
259    } else {
260        quote!(trillium::Handler::upgrade(&#handler, upgrade))
261    };
262
263    quote! {
264        impl #impl_generics trillium::Handler for #struct_name #ty_generics #where_clause {
265            async fn run(&self, conn: trillium::Conn) -> trillium::Conn { #run.await }
266            async fn init(&mut self, info: &mut trillium::Info) { #init.await; }
267            async fn before_send(&self, conn: trillium::Conn) -> trillium::Conn { #before_send.await }
268            fn name(&self) -> std::borrow::Cow<'static, str> { #name }
269            fn has_upgrade(&self, upgrade: &trillium::Upgrade) -> bool { #has_upgrade }
270            async fn upgrade(&self, upgrade: trillium::Upgrade) { #upgrade.await }
271        }
272    }
273    .into()
274}