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}