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, 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 let Some(m) = try_extract_wrapper(ty, JSON) {
123 let output = m
124 .first_type()
125 .ok_or_else(|| Error::empty_generic_args(ty.span(), JSON))?
126 .clone();
127 return Ok((output, None, false));
128 }
129
130 if let Some(result_match) = try_extract_wrapper(ty, RESULT) {
132 let first = result_match
133 .first_type()
134 .ok_or_else(|| Error::empty_generic_args(ty.span(), RESULT))?;
135
136 let json_match = try_extract_wrapper(first, JSON).ok_or_else(|| {
137 Error::invalid_handler_sig(
138 first.span(),
139 fn_name,
140 "Result's first type argument must be Json<T>",
141 )
142 })?;
143
144 let output = json_match
145 .first_type()
146 .ok_or_else(|| Error::empty_generic_args(first.span(), JSON))?
147 .clone();
148
149 let error_type = result_match.second_type().cloned();
150 return Ok((output, error_type, false));
151 }
152
153 Err(Error::invalid_handler_sig(
154 ty.span(),
155 fn_name,
156 "return type must be Json<T>, Result<Json<T>, E>, or Sse<impl Stream<...>>",
157 ))
158}
159
160fn extract_state_param(
162 inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
163) -> Option<Type> {
164 for arg in inputs {
165 if let FnArg::Typed(pat_type) = arg
166 && let Some(m) = try_extract_wrapper(&pat_type.ty, STATE)
167 {
168 return m.first_type().cloned();
169 }
170 }
171 None
172}
173
174fn extract_json_param(inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>) -> Option<Type> {
176 for arg in inputs {
177 if let FnArg::Typed(pat_type) = arg
178 && let Some(m) = try_extract_wrapper(&pat_type.ty, JSON)
179 {
180 return m.first_type().cloned();
181 }
182 }
183 None
184}
185
186fn extract_query_param(
188 inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
189) -> Option<Type> {
190 for arg in inputs {
191 if let FnArg::Typed(pat_type) = arg
192 && let Some(m) = try_extract_wrapper(&pat_type.ty, QUERY)
193 {
194 return m.first_type().cloned();
195 }
196 }
197 None
198}
199
200fn extract_path_params(
205 inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
206) -> Vec<(String, Type)> {
207 let mut params = Vec::new();
208 for arg in inputs {
209 if let FnArg::Typed(pat_type) = arg
210 && let Some(m) = try_extract_wrapper(&pat_type.ty, PATH)
211 && let Some(inner) = m.first_type()
212 {
213 let name = extract_path_binding_name(&pat_type.pat);
214 params.push((name, inner.clone()));
215 }
216 }
217 params
218}
219
220fn extract_path_binding_name(pat: &syn::Pat) -> String {
224 if let syn::Pat::TupleStruct(ts) = pat
226 && let Some(first) = ts.elems.first()
227 && let syn::Pat::Ident(id) = first
228 {
229 return id.ident.to_string();
230 }
231 if let syn::Pat::Ident(id) = pat {
233 return id.ident.to_string();
234 }
235 "param".to_string()
236}
237
238#[cfg(test)]
243mod tests {
244 use super::*;
245 use crate::errors::type_display;
246 use syn::parse_quote;
247
248 fn sig(func: ItemFn) -> HandlerSignature {
249 extract_handler_signature(&func).unwrap()
250 }
251
252 fn sig_err(func: ItemFn) -> Error {
253 extract_handler_signature(&func).unwrap_err()
254 }
255
256 #[test]
259 fn json_return_only() {
260 let f: ItemFn = parse_quote! {
261 async fn handler() -> Json<Planet> {}
262 };
263 let s = sig(f);
264 assert_eq!(s.fn_name, "handler");
265 assert_eq!(type_display(&s.output_type), "Planet");
266 assert!(s.error_type.is_none());
267 assert!(s.input_type.is_none());
268 assert!(s.state_type.is_none());
269 assert!(s.is_async);
270 assert!(!s.is_streaming);
271 }
272
273 #[test]
274 fn result_json_return() {
275 let f: ItemFn = parse_quote! {
276 async fn handler() -> Result<Json<Planet>, AppError> {}
277 };
278 let s = sig(f);
279 assert_eq!(type_display(&s.output_type), "Planet");
280 assert_eq!(type_display(s.error_type.as_ref().unwrap()), "AppError");
281 }
282
283 #[test]
284 fn qualified_result_return() {
285 let f: ItemFn = parse_quote! {
286 async fn handler() -> std::result::Result<Json<Planet>, AppError> {}
287 };
288 let s = sig(f);
289 assert_eq!(type_display(&s.output_type), "Planet");
290 }
291
292 #[test]
293 fn qualified_json_return() {
294 let f: ItemFn = parse_quote! {
295 async fn handler() -> axum::extract::Json<Planet> {}
296 };
297 let s = sig(f);
298 assert_eq!(type_display(&s.output_type), "Planet");
299 }
300
301 #[test]
302 fn state_param_extracted() {
303 let f: ItemFn = parse_quote! {
304 async fn handler(State(db): State<Db>) -> Json<Planet> {}
305 };
306 let s = sig(f);
307 assert_eq!(type_display(s.state_type.as_ref().unwrap()), "Db");
308 }
309
310 #[test]
311 fn json_param_extracted() {
312 let f: ItemFn = parse_quote! {
313 async fn handler(Json(body): Json<CreatePlanet>) -> Json<Planet> {}
314 };
315 let s = sig(f);
316 assert_eq!(type_display(s.input_type.as_ref().unwrap()), "CreatePlanet");
317 }
318
319 #[test]
320 fn both_state_and_json_params() {
321 let f: ItemFn = parse_quote! {
322 async fn handler(State(db): State<Db>, Json(body): Json<CreatePlanet>) -> Result<Json<Planet>, AppError> {}
323 };
324 let s = sig(f);
325 assert_eq!(type_display(s.state_type.as_ref().unwrap()), "Db");
326 assert_eq!(type_display(s.input_type.as_ref().unwrap()), "CreatePlanet");
327 assert_eq!(type_display(&s.output_type), "Planet");
328 assert!(s.error_type.is_some());
329 }
330
331 #[test]
332 fn path_param_extracted() {
333 let f: ItemFn = parse_quote! {
334 async fn handler(Path(id): Path<i32>) -> Json<Planet> {}
335 };
336 let s = sig(f);
337 assert_eq!(s.path_params.len(), 1);
338 assert_eq!(s.path_params[0].0, "id");
339 assert_eq!(type_display(&s.path_params[0].1), "i32");
340 }
341
342 #[test]
343 fn multiple_path_params_extracted() {
344 let f: ItemFn = parse_quote! {
345 async fn handler(Path(ws_id): Path<i32>, Path(proj_id): Path<String>) -> Json<Planet> {}
346 };
347 let s = sig(f);
348 assert_eq!(s.path_params.len(), 2);
349 assert_eq!(s.path_params[0].0, "ws_id");
350 assert_eq!(type_display(&s.path_params[0].1), "i32");
351 assert_eq!(s.path_params[1].0, "proj_id");
352 assert_eq!(type_display(&s.path_params[1].1), "String");
353 }
354
355 #[test]
356 fn path_and_query_params_both_extracted() {
357 let f: ItemFn = parse_quote! {
358 async fn handler(Path(id): Path<i32>, Query(q): Query<SearchQuery>) -> Json<Planet> {}
359 };
360 let s = sig(f);
361 assert_eq!(s.path_params.len(), 1);
362 assert_eq!(s.path_params[0].0, "id");
363 assert_eq!(type_display(s.query_type.as_ref().unwrap()), "SearchQuery");
364 }
365
366 #[test]
367 fn no_path_params_gives_empty_vec() {
368 let f: ItemFn = parse_quote! {
369 async fn handler(State(db): State<Db>) -> Json<Planet> {}
370 };
371 let s = sig(f);
372 assert!(s.path_params.is_empty());
373 }
374
375 #[test]
376 fn sync_function_allowed() {
377 let f: ItemFn = parse_quote! {
378 fn handler() -> Json<Planet> {}
379 };
380 let s = sig(f);
381 assert!(!s.is_async);
382 }
383
384 #[test]
385 fn sse_return_is_streaming() {
386 let f: ItemFn = parse_quote! {
387 async fn stream_events(State(_state): State<AppState>) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {}
388 };
389 let s = sig(f);
390 assert!(s.is_streaming);
391 assert!(s.error_type.is_none());
392 assert_eq!(type_display(&s.output_type), "()");
394 }
395
396 #[test]
397 fn qualified_sse_return_is_streaming() {
398 let f: ItemFn = parse_quote! {
399 async fn handler() -> axum::response::Sse<SomeStream> {}
400 };
401 let s = sig(f);
402 assert!(s.is_streaming);
403 }
404
405 #[test]
408 fn no_return_type_error() {
409 let f: ItemFn = parse_quote! {
410 async fn handler() {}
411 };
412 let err = sig_err(f);
413 assert!(err.to_string().contains("has no return type"));
414 }
415
416 #[test]
417 fn bare_type_return_error() {
418 let f: ItemFn = parse_quote! {
419 async fn handler() -> Vec<Planet> {}
420 };
421 let err = sig_err(f);
422 assert!(err.to_string().contains("return type must be Json<T>"));
423 }
424
425 #[test]
426 fn result_without_json_inner_error() {
427 let f: ItemFn = parse_quote! {
428 async fn handler() -> Result<Vec<Planet>, AppError> {}
429 };
430 let err = sig_err(f);
431 assert!(
432 err.to_string()
433 .contains("Result's first type argument must be Json<T>")
434 );
435 }
436}