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, 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    /// The resolved output type:
36    /// - `Json<T>` return → `T`
37    /// - `Result<Json<T>, E>` return → `T`
38    /// - `Sse<...>` return → unit `()` (output type comes from `stream_event` attribute)
39    pub output_type: Type,
40    /// The `E` in `Result<_, E>`, if present.
41    pub error_type: Option<Type>,
42    /// Whether the handler returns `Sse<...>` (an SSE streaming response).
43    pub is_streaming: bool,
44    /// Whether the function is declared `async`.
45    pub is_async: bool,
46}
47
48// ---------------------------------------------------------------------------
49// Public API
50// ---------------------------------------------------------------------------
51
52/// Extract a [`HandlerSignature`] from a handler function.
53///
54/// Validates:
55/// - The return type is `Json<T>` or `Result<Json<T>, E>` (not a bare type)
56/// - Collects `State<S>`, `Json<T>`, and `Query<T>` parameters when present
57///
58/// Handlers that return neither `Json<T>` nor `Result<Json<T>, E>` are
59/// rejected with an error pointing at the return type token.
60pub 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
83// ---------------------------------------------------------------------------
84// Internal extraction helpers
85// ---------------------------------------------------------------------------
86
87/// Extract the unwrapped output type, optional error type, and streaming flag
88/// from a return type.
89///
90/// Accepts:
91/// - `-> Json<T>` → (T, None, false)
92/// - `-> Result<Json<T>, E>` → (T, Some(E), false)
93/// - `-> Sse<...>` → ((), None, true)
94fn 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    // Case 1: Sse<...> — streaming handler; output type comes from stream_event attribute
109    if try_extract_wrapper(ty, SSE).is_some() {
110        let unit: Type = syn::parse_quote! { () };
111        return Ok((unit, None, true));
112    }
113
114    // Case 2: Json<T>
115    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    // Case 3: Result<Json<T>, E>
124    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
153/// Find a `State<S>` parameter and return the inner `S`.
154fn 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
167/// Find the first `Json<T>` parameter and return the inner `T`.
168fn 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
179/// Find the first `Query<T>` parameter and return the inner `T`.
180fn 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// ---------------------------------------------------------------------------
194// Tests
195// ---------------------------------------------------------------------------
196
197#[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    // --- valid signatures ---
212
213    #[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        // output_type is unit () for SSE handlers
304        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    // --- invalid signatures ---
317
318    #[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}