Skip to main content

rorpc_parse/
functions.rs

1//! Handler function signature analysis.
2//!
3//! Extracts a fully-typed [`HandlerSignature`] from a `syn::ItemFn` using
4//! AST-based type inspection via [`crate::types`]. All type matching is done
5//! on path segment idents — never on string representations.
6
7use 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// ---------------------------------------------------------------------------
16// HandlerSignature
17// ---------------------------------------------------------------------------
18
19/// Fully analysed handler function signature.
20///
21/// Produced by [`extract_handler_signature`]. All fields are resolved against
22/// the actual AST — no string-based type inference.
23#[derive(Debug)]
24pub struct HandlerSignature {
25    /// The function's identifier, e.g. `"list_planets"`.
26    pub fn_name: String,
27    /// Span of the function identifier for error reporting.
28    pub fn_span: Span,
29    /// The `S` in a `State<S>` parameter, if present.
30    pub state_type: Option<Type>,
31    /// The `T` in a `Json<T>` parameter, if present.
32    pub input_type: Option<Type>,
33    /// The `T` in a `Query<T>` parameter, if present.
34    pub query_type: Option<Type>,
35    /// Ordered list of `(binding_name, rust_type)` pairs from `Path<T>` parameters.
36    ///
37    /// E.g., `Path(id): Path<i32>` → `[("id", Type::i32)]`
38    /// Multiple params: `Path(id): Path<i32>, Path(slug): Path<String>` → `[("id", i32), ("slug", String)]`
39    pub path_params: Vec<(String, Type)>,
40    /// The resolved output type:
41    /// - `Json<T>` return → `T`
42    /// - `Result<Json<T>, E>` return → `T`
43    /// - `Sse<...>` return → unit `()` (output type comes from `data` attribute)
44    pub output_type: Type,
45    /// The `E` in `Result<_, E>`, if present.
46    pub error_type: Option<Type>,
47    /// Whether the handler returns `Sse<...>` (an SSE streaming response).
48    pub is_streaming: bool,
49    /// Whether the function is declared `async`.
50    pub is_async: bool,
51}
52
53// ---------------------------------------------------------------------------
54// Public API
55// ---------------------------------------------------------------------------
56
57/// Extract a [`HandlerSignature`] from a handler function.
58///
59/// Validates:
60/// - The return type is `Json<T>` or `Result<Json<T>, E>` (not a bare type)
61/// - Collects `State<S>`, `Json<T>`, and `Query<T>` parameters when present
62///
63/// Handlers that return neither `Json<T>` nor `Result<Json<T>, E>` are
64/// rejected with an error pointing at the return type token.
65pub 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
90// ---------------------------------------------------------------------------
91// Internal extraction helpers
92// ---------------------------------------------------------------------------
93
94/// Extract the unwrapped output type, optional error type, and streaming flag
95/// from a return type.
96///
97/// Accepts:
98/// - `-> Json<T>` → (T, None, false)
99/// - `-> Result<Json<T>, E>` → (T, Some(E), false)
100/// - `-> Sse<...>` → ((), None, true)
101fn 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    // Case 1: Sse<...> — streaming handler; output type comes from data attribute
116    if try_extract_wrapper(ty, SSE).is_some() {
117        let unit: Type = syn::parse_quote! { () };
118        return Ok((unit, None, true));
119    }
120
121    // Case 2: Json<T>
122    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    // Case 3: Result<Json<T>, E>
131    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
160/// Find a `State<S>` parameter and return the inner `S`.
161fn 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
174/// Find the first `Json<T>` parameter and return the inner `T`.
175fn 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
186/// Find the first `Query<T>` parameter and return the inner `T`.
187fn 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
200/// Collect all `Path<T>` parameters, returning `(binding_name, inner_type)` pairs.
201///
202/// `Path(id): Path<i32>` → `("id", Type::i32)`
203/// Preserves declaration order so types can be zipped with path template params.
204fn 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
220/// Extract the binding name from a `Path(name)` or `Path { name }` pattern.
221///
222/// `Path(id)` → `"id"`, falls back to `"param"` for unrecognised patterns.
223fn extract_path_binding_name(pat: &syn::Pat) -> String {
224    // Most common form: Path(id) — a TupleStruct pattern
225    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    // Plain ident (uncommon but valid)
232    if let syn::Pat::Ident(id) = pat {
233        return id.ident.to_string();
234    }
235    "param".to_string()
236}
237
238// ---------------------------------------------------------------------------
239// Tests
240// ---------------------------------------------------------------------------
241
242#[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    // --- valid signatures ---
257
258    #[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        // output_type is unit () for SSE handlers
393        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    // --- invalid signatures ---
406
407    #[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}