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