1use proc_macro2::Span;
8use syn::{FnArg, ItemFn, ReturnType, Type, spanned::Spanned};
9
10use crate::{
11 errors::{Error, Result},
12 types::{JSON, QUERY, RESULT, SSE, STATE, try_extract_wrapper},
13};
14
15#[derive(Debug)]
24pub struct HandlerSignature {
25 pub fn_name: String,
27 pub fn_span: Span,
29 pub state_type: Option<Type>,
31 pub input_type: Option<Type>,
33 pub query_type: Option<Type>,
35 pub output_type: Type,
40 pub error_type: Option<Type>,
42 pub is_streaming: bool,
44 pub is_async: bool,
46}
47
48pub fn extract_handler_signature(func: &ItemFn) -> Result<HandlerSignature> {
61 let fn_name = func.sig.ident.to_string();
62 let fn_span = func.sig.ident.span();
63 let is_async = func.sig.asyncness.is_some();
64
65 let (output_type, error_type, is_streaming) = extract_return_types(&func.sig.output, &fn_name)?;
66 let state_type = extract_state_param(&func.sig.inputs);
67 let input_type = extract_json_param(&func.sig.inputs);
68 let query_type = extract_query_param(&func.sig.inputs);
69
70 Ok(HandlerSignature {
71 fn_name,
72 fn_span,
73 state_type,
74 input_type,
75 query_type,
76 output_type,
77 error_type,
78 is_streaming,
79 is_async,
80 })
81}
82
83fn extract_return_types(
95 return_type: &ReturnType,
96 fn_name: &str,
97) -> Result<(Type, Option<Type>, bool)> {
98 let ty = match return_type {
99 ReturnType::Default => {
100 return Err(Error::missing_return_type(
101 proc_macro2::Span::call_site(),
102 fn_name,
103 ));
104 }
105 ReturnType::Type(_, ty) => ty.as_ref(),
106 };
107
108 if try_extract_wrapper(ty, SSE).is_some() {
110 let unit: Type = syn::parse_quote! { () };
111 return Ok((unit, None, true));
112 }
113
114 if let Some(m) = try_extract_wrapper(ty, JSON) {
116 let output = m
117 .first_type()
118 .ok_or_else(|| Error::empty_generic_args(ty.span(), JSON))?
119 .clone();
120 return Ok((output, None, false));
121 }
122
123 if let Some(result_match) = try_extract_wrapper(ty, RESULT) {
125 let first = result_match
126 .first_type()
127 .ok_or_else(|| Error::empty_generic_args(ty.span(), RESULT))?;
128
129 let json_match = try_extract_wrapper(first, JSON).ok_or_else(|| {
130 Error::invalid_handler_sig(
131 first.span(),
132 fn_name,
133 "Result's first type argument must be Json<T>",
134 )
135 })?;
136
137 let output = json_match
138 .first_type()
139 .ok_or_else(|| Error::empty_generic_args(first.span(), JSON))?
140 .clone();
141
142 let error_type = result_match.second_type().cloned();
143 return Ok((output, error_type, false));
144 }
145
146 Err(Error::invalid_handler_sig(
147 ty.span(),
148 fn_name,
149 "return type must be Json<T>, Result<Json<T>, E>, or Sse<impl Stream<...>>",
150 ))
151}
152
153fn extract_state_param(
155 inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
156) -> Option<Type> {
157 for arg in inputs {
158 if let FnArg::Typed(pat_type) = arg
159 && let Some(m) = try_extract_wrapper(&pat_type.ty, STATE)
160 {
161 return m.first_type().cloned();
162 }
163 }
164 None
165}
166
167fn extract_json_param(inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>) -> Option<Type> {
169 for arg in inputs {
170 if let FnArg::Typed(pat_type) = arg
171 && let Some(m) = try_extract_wrapper(&pat_type.ty, JSON)
172 {
173 return m.first_type().cloned();
174 }
175 }
176 None
177}
178
179fn extract_query_param(
181 inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
182) -> Option<Type> {
183 for arg in inputs {
184 if let FnArg::Typed(pat_type) = arg
185 && let Some(m) = try_extract_wrapper(&pat_type.ty, QUERY)
186 {
187 return m.first_type().cloned();
188 }
189 }
190 None
191}
192
193#[cfg(test)]
198mod tests {
199 use super::*;
200 use crate::errors::type_display;
201 use syn::parse_quote;
202
203 fn sig(func: ItemFn) -> HandlerSignature {
204 extract_handler_signature(&func).unwrap()
205 }
206
207 fn sig_err(func: ItemFn) -> Error {
208 extract_handler_signature(&func).unwrap_err()
209 }
210
211 #[test]
214 fn json_return_only() {
215 let f: ItemFn = parse_quote! {
216 async fn handler() -> Json<Planet> {}
217 };
218 let s = sig(f);
219 assert_eq!(s.fn_name, "handler");
220 assert_eq!(type_display(&s.output_type), "Planet");
221 assert!(s.error_type.is_none());
222 assert!(s.input_type.is_none());
223 assert!(s.state_type.is_none());
224 assert!(s.is_async);
225 assert!(!s.is_streaming);
226 }
227
228 #[test]
229 fn result_json_return() {
230 let f: ItemFn = parse_quote! {
231 async fn handler() -> Result<Json<Planet>, AppError> {}
232 };
233 let s = sig(f);
234 assert_eq!(type_display(&s.output_type), "Planet");
235 assert_eq!(type_display(s.error_type.as_ref().unwrap()), "AppError");
236 }
237
238 #[test]
239 fn qualified_result_return() {
240 let f: ItemFn = parse_quote! {
241 async fn handler() -> std::result::Result<Json<Planet>, AppError> {}
242 };
243 let s = sig(f);
244 assert_eq!(type_display(&s.output_type), "Planet");
245 }
246
247 #[test]
248 fn qualified_json_return() {
249 let f: ItemFn = parse_quote! {
250 async fn handler() -> axum::extract::Json<Planet> {}
251 };
252 let s = sig(f);
253 assert_eq!(type_display(&s.output_type), "Planet");
254 }
255
256 #[test]
257 fn state_param_extracted() {
258 let f: ItemFn = parse_quote! {
259 async fn handler(State(db): State<Db>) -> Json<Planet> {}
260 };
261 let s = sig(f);
262 assert_eq!(type_display(s.state_type.as_ref().unwrap()), "Db");
263 }
264
265 #[test]
266 fn json_param_extracted() {
267 let f: ItemFn = parse_quote! {
268 async fn handler(Json(body): Json<CreatePlanet>) -> Json<Planet> {}
269 };
270 let s = sig(f);
271 assert_eq!(type_display(s.input_type.as_ref().unwrap()), "CreatePlanet");
272 }
273
274 #[test]
275 fn both_state_and_json_params() {
276 let f: ItemFn = parse_quote! {
277 async fn handler(State(db): State<Db>, Json(body): Json<CreatePlanet>) -> Result<Json<Planet>, AppError> {}
278 };
279 let s = sig(f);
280 assert_eq!(type_display(s.state_type.as_ref().unwrap()), "Db");
281 assert_eq!(type_display(s.input_type.as_ref().unwrap()), "CreatePlanet");
282 assert_eq!(type_display(&s.output_type), "Planet");
283 assert!(s.error_type.is_some());
284 }
285
286 #[test]
287 fn sync_function_allowed() {
288 let f: ItemFn = parse_quote! {
289 fn handler() -> Json<Planet> {}
290 };
291 let s = sig(f);
292 assert!(!s.is_async);
293 }
294
295 #[test]
296 fn sse_return_is_streaming() {
297 let f: ItemFn = parse_quote! {
298 async fn stream_events(State(_state): State<AppState>) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {}
299 };
300 let s = sig(f);
301 assert!(s.is_streaming);
302 assert!(s.error_type.is_none());
303 assert_eq!(type_display(&s.output_type), "()");
305 }
306
307 #[test]
308 fn qualified_sse_return_is_streaming() {
309 let f: ItemFn = parse_quote! {
310 async fn handler() -> axum::response::Sse<SomeStream> {}
311 };
312 let s = sig(f);
313 assert!(s.is_streaming);
314 }
315
316 #[test]
319 fn no_return_type_error() {
320 let f: ItemFn = parse_quote! {
321 async fn handler() {}
322 };
323 let err = sig_err(f);
324 assert!(err.to_string().contains("has no return type"));
325 }
326
327 #[test]
328 fn bare_type_return_error() {
329 let f: ItemFn = parse_quote! {
330 async fn handler() -> Vec<Planet> {}
331 };
332 let err = sig_err(f);
333 assert!(err.to_string().contains("return type must be Json<T>"));
334 }
335
336 #[test]
337 fn result_without_json_inner_error() {
338 let f: ItemFn = parse_quote! {
339 async fn handler() -> Result<Vec<Planet>, AppError> {}
340 };
341 let err = sig_err(f);
342 assert!(
343 err.to_string()
344 .contains("Result's first type argument must be Json<T>")
345 );
346 }
347}