1use proc_macro2::Span;
8use syn::{FnArg, ItemFn, ReturnType, Type, spanned::Spanned};
9
10use crate::{
11 errors::{Error, Result},
12 types::{JSON, PATH, QUERY, RESULT, SSE, STATE, STATUSCODE, last_ident_matches, 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 path_params: Vec<(String, Type)>,
40 pub output_type: Type,
45 pub error_type: Option<Type>,
47 pub is_streaming: bool,
49 pub is_async: bool,
51}
52
53pub fn extract_handler_signature(func: &ItemFn) -> Result<HandlerSignature> {
66 let fn_name = func.sig.ident.to_string();
67 let fn_span = func.sig.ident.span();
68 let is_async = func.sig.asyncness.is_some();
69
70 let (output_type, error_type, is_streaming) = extract_return_types(&func.sig.output, &fn_name)?;
71 let state_type = extract_state_param(&func.sig.inputs);
72 let input_type = extract_json_param(&func.sig.inputs);
73 let query_type = extract_query_param(&func.sig.inputs);
74 let path_params = extract_path_params(&func.sig.inputs);
75
76 Ok(HandlerSignature {
77 fn_name,
78 fn_span,
79 state_type,
80 input_type,
81 query_type,
82 path_params,
83 output_type,
84 error_type,
85 is_streaming,
86 is_async,
87 })
88}
89
90fn extract_return_types(
102 return_type: &ReturnType,
103 fn_name: &str,
104) -> Result<(Type, Option<Type>, bool)> {
105 let ty = match return_type {
106 ReturnType::Default => {
107 return Err(Error::missing_return_type(
108 proc_macro2::Span::call_site(),
109 fn_name,
110 ));
111 }
112 ReturnType::Type(_, ty) => ty.as_ref(),
113 };
114
115 if try_extract_wrapper(ty, SSE).is_some() {
117 let unit: Type = syn::parse_quote! { () };
118 return Ok((unit, None, true));
119 }
120
121 if last_ident_matches(ty, STATUSCODE) {
123 let unit: Type = syn::parse_quote! { () };
124 return Ok((unit, None, false));
125 }
126
127 if let Some(m) = try_extract_wrapper(ty, JSON) {
129 let output = m
130 .first_type()
131 .ok_or_else(|| Error::empty_generic_args(ty.span(), JSON))?
132 .clone();
133 return Ok((output, None, false));
134 }
135
136 if let Some(result_match) = try_extract_wrapper(ty, RESULT) {
138 let first = result_match
139 .first_type()
140 .ok_or_else(|| Error::empty_generic_args(ty.span(), RESULT))?;
141
142 let json_match = try_extract_wrapper(first, JSON).ok_or_else(|| {
143 Error::invalid_handler_sig(
144 first.span(),
145 fn_name,
146 "Result's first type argument must be Json<T>",
147 )
148 })?;
149
150 let output = json_match
151 .first_type()
152 .ok_or_else(|| Error::empty_generic_args(first.span(), JSON))?
153 .clone();
154
155 let error_type = result_match.second_type().cloned();
156 return Ok((output, error_type, false));
157 }
158
159 Err(Error::invalid_handler_sig(
160 ty.span(),
161 fn_name,
162 "return type must be Json<T>, Result<Json<T>, E>, StatusCode, or Sse<impl Stream<...>>",
163 ))
164}
165
166fn extract_state_param(
168 inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
169) -> Option<Type> {
170 for arg in inputs {
171 if let FnArg::Typed(pat_type) = arg
172 && let Some(m) = try_extract_wrapper(&pat_type.ty, STATE)
173 {
174 return m.first_type().cloned();
175 }
176 }
177 None
178}
179
180fn extract_json_param(inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>) -> Option<Type> {
182 for arg in inputs {
183 if let FnArg::Typed(pat_type) = arg
184 && let Some(m) = try_extract_wrapper(&pat_type.ty, JSON)
185 {
186 return m.first_type().cloned();
187 }
188 }
189 None
190}
191
192fn extract_query_param(
194 inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
195) -> Option<Type> {
196 for arg in inputs {
197 if let FnArg::Typed(pat_type) = arg
198 && let Some(m) = try_extract_wrapper(&pat_type.ty, QUERY)
199 {
200 return m.first_type().cloned();
201 }
202 }
203 None
204}
205
206fn extract_path_params(
211 inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
212) -> Vec<(String, Type)> {
213 let mut params = Vec::new();
214 for arg in inputs {
215 if let FnArg::Typed(pat_type) = arg
216 && let Some(m) = try_extract_wrapper(&pat_type.ty, PATH)
217 && let Some(inner) = m.first_type()
218 {
219 let name = extract_path_binding_name(&pat_type.pat);
220 params.push((name, inner.clone()));
221 }
222 }
223 params
224}
225
226fn extract_path_binding_name(pat: &syn::Pat) -> String {
230 if let syn::Pat::TupleStruct(ts) = pat
232 && let Some(first) = ts.elems.first()
233 && let syn::Pat::Ident(id) = first
234 {
235 return id.ident.to_string();
236 }
237 if let syn::Pat::Ident(id) = pat {
239 return id.ident.to_string();
240 }
241 "param".to_string()
242}
243
244#[cfg(test)]
249mod tests {
250 use super::*;
251 use crate::errors::type_display;
252 use syn::parse_quote;
253
254 fn sig(func: ItemFn) -> HandlerSignature {
255 extract_handler_signature(&func).unwrap()
256 }
257
258 fn sig_err(func: ItemFn) -> Error {
259 extract_handler_signature(&func).unwrap_err()
260 }
261
262 #[test]
265 fn json_return_only() {
266 let f: ItemFn = parse_quote! {
267 async fn handler() -> Json<Planet> {}
268 };
269 let s = sig(f);
270 assert_eq!(s.fn_name, "handler");
271 assert_eq!(type_display(&s.output_type), "Planet");
272 assert!(s.error_type.is_none());
273 assert!(s.input_type.is_none());
274 assert!(s.state_type.is_none());
275 assert!(s.is_async);
276 assert!(!s.is_streaming);
277 }
278
279 #[test]
280 fn result_json_return() {
281 let f: ItemFn = parse_quote! {
282 async fn handler() -> Result<Json<Planet>, AppError> {}
283 };
284 let s = sig(f);
285 assert_eq!(type_display(&s.output_type), "Planet");
286 assert_eq!(type_display(s.error_type.as_ref().unwrap()), "AppError");
287 }
288
289 #[test]
290 fn qualified_result_return() {
291 let f: ItemFn = parse_quote! {
292 async fn handler() -> std::result::Result<Json<Planet>, AppError> {}
293 };
294 let s = sig(f);
295 assert_eq!(type_display(&s.output_type), "Planet");
296 }
297
298 #[test]
299 fn qualified_json_return() {
300 let f: ItemFn = parse_quote! {
301 async fn handler() -> axum::extract::Json<Planet> {}
302 };
303 let s = sig(f);
304 assert_eq!(type_display(&s.output_type), "Planet");
305 }
306
307 #[test]
308 fn state_param_extracted() {
309 let f: ItemFn = parse_quote! {
310 async fn handler(State(db): State<Db>) -> Json<Planet> {}
311 };
312 let s = sig(f);
313 assert_eq!(type_display(s.state_type.as_ref().unwrap()), "Db");
314 }
315
316 #[test]
317 fn json_param_extracted() {
318 let f: ItemFn = parse_quote! {
319 async fn handler(Json(body): Json<CreatePlanet>) -> Json<Planet> {}
320 };
321 let s = sig(f);
322 assert_eq!(type_display(s.input_type.as_ref().unwrap()), "CreatePlanet");
323 }
324
325 #[test]
326 fn both_state_and_json_params() {
327 let f: ItemFn = parse_quote! {
328 async fn handler(State(db): State<Db>, Json(body): Json<CreatePlanet>) -> Result<Json<Planet>, AppError> {}
329 };
330 let s = sig(f);
331 assert_eq!(type_display(s.state_type.as_ref().unwrap()), "Db");
332 assert_eq!(type_display(s.input_type.as_ref().unwrap()), "CreatePlanet");
333 assert_eq!(type_display(&s.output_type), "Planet");
334 assert!(s.error_type.is_some());
335 }
336
337 #[test]
338 fn path_param_extracted() {
339 let f: ItemFn = parse_quote! {
340 async fn handler(Path(id): Path<i32>) -> Json<Planet> {}
341 };
342 let s = sig(f);
343 assert_eq!(s.path_params.len(), 1);
344 assert_eq!(s.path_params[0].0, "id");
345 assert_eq!(type_display(&s.path_params[0].1), "i32");
346 }
347
348 #[test]
349 fn multiple_path_params_extracted() {
350 let f: ItemFn = parse_quote! {
351 async fn handler(Path(ws_id): Path<i32>, Path(proj_id): Path<String>) -> Json<Planet> {}
352 };
353 let s = sig(f);
354 assert_eq!(s.path_params.len(), 2);
355 assert_eq!(s.path_params[0].0, "ws_id");
356 assert_eq!(type_display(&s.path_params[0].1), "i32");
357 assert_eq!(s.path_params[1].0, "proj_id");
358 assert_eq!(type_display(&s.path_params[1].1), "String");
359 }
360
361 #[test]
362 fn path_and_query_params_both_extracted() {
363 let f: ItemFn = parse_quote! {
364 async fn handler(Path(id): Path<i32>, Query(q): Query<SearchQuery>) -> Json<Planet> {}
365 };
366 let s = sig(f);
367 assert_eq!(s.path_params.len(), 1);
368 assert_eq!(s.path_params[0].0, "id");
369 assert_eq!(type_display(s.query_type.as_ref().unwrap()), "SearchQuery");
370 }
371
372 #[test]
373 fn no_path_params_gives_empty_vec() {
374 let f: ItemFn = parse_quote! {
375 async fn handler(State(db): State<Db>) -> Json<Planet> {}
376 };
377 let s = sig(f);
378 assert!(s.path_params.is_empty());
379 }
380
381 #[test]
382 fn sync_function_allowed() {
383 let f: ItemFn = parse_quote! {
384 fn handler() -> Json<Planet> {}
385 };
386 let s = sig(f);
387 assert!(!s.is_async);
388 }
389
390 #[test]
391 fn sse_return_is_streaming() {
392 let f: ItemFn = parse_quote! {
393 async fn stream_events(State(_state): State<AppState>) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {}
394 };
395 let s = sig(f);
396 assert!(s.is_streaming);
397 assert!(s.error_type.is_none());
398 assert_eq!(type_display(&s.output_type), "()");
400 }
401
402 #[test]
403 fn qualified_sse_return_is_streaming() {
404 let f: ItemFn = parse_quote! {
405 async fn handler() -> axum::response::Sse<SomeStream> {}
406 };
407 let s = sig(f);
408 assert!(s.is_streaming);
409 }
410
411 #[test]
412 fn statuscode_return_maps_to_unit() {
413 let f: ItemFn = parse_quote! {
414 async fn handler() -> StatusCode {}
415 };
416 let s = sig(f);
417 assert_eq!(type_display(&s.output_type), "()");
418 assert!(!s.is_streaming);
419 }
420
421 #[test]
422 fn qualified_statuscode_works() {
423 let f: ItemFn = parse_quote! {
424 async fn handler() -> http::StatusCode {}
425 };
426 let s = sig(f);
427 assert_eq!(type_display(&s.output_type), "()");
428 }
429
430 #[test]
433 fn no_return_type_error() {
434 let f: ItemFn = parse_quote! {
435 async fn handler() {}
436 };
437 let err = sig_err(f);
438 assert!(err.to_string().contains("has no return type"));
439 }
440
441 #[test]
442 fn bare_type_return_error() {
443 let f: ItemFn = parse_quote! {
444 async fn handler() -> Vec<Planet> {}
445 };
446 let err = sig_err(f);
447 assert!(err.to_string().contains("return type must be Json<T>"));
448 }
449
450 #[test]
451 fn result_without_json_inner_error() {
452 let f: ItemFn = parse_quote! {
453 async fn handler() -> Result<Vec<Planet>, AppError> {}
454 };
455 let err = sig_err(f);
456 assert!(
457 err.to_string()
458 .contains("Result's first type argument must be Json<T>")
459 );
460 }
461}