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, 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 resolved output type:
34    /// - `Json<T>` return → `T`
35    /// - `Result<Json<T>, E>` return → `T`
36    /// - `Sse<...>` return → unit `()` (output type comes from `stream_event` attribute)
37    pub output_type: Type,
38    /// The `E` in `Result<_, E>`, if present.
39    pub error_type: Option<Type>,
40    /// Whether the handler returns `Sse<...>` (an SSE streaming response).
41    pub is_streaming: bool,
42    /// Whether the function is declared `async`.
43    pub is_async: bool,
44}
45
46// ---------------------------------------------------------------------------
47// Public API
48// ---------------------------------------------------------------------------
49
50/// Extract a [`HandlerSignature`] from a handler function.
51///
52/// Validates:
53/// - The return type is `Json<T>` or `Result<Json<T>, E>` (not a bare type)
54/// - Collects `State<S>` and `Json<T>` parameters when present
55///
56/// Handlers that return neither `Json<T>` nor `Result<Json<T>, E>` are
57/// rejected with an error pointing at the return type token.
58pub 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
79// ---------------------------------------------------------------------------
80// Internal extraction helpers
81// ---------------------------------------------------------------------------
82
83/// Extract the unwrapped output type, optional error type, and streaming flag
84/// from a return type.
85///
86/// Accepts:
87/// - `-> Json<T>` → (T, None, false)
88/// - `-> Result<Json<T>, E>` → (T, Some(E), false)
89/// - `-> Sse<...>` → ((), None, true)
90fn 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    // Case 1: Sse<...> — streaming handler; output type comes from stream_event attribute
105    if try_extract_wrapper(ty, SSE).is_some() {
106        let unit: Type = syn::parse_quote! { () };
107        return Ok((unit, None, true));
108    }
109
110    // Case 2: Json<T>
111    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    // Case 3: Result<Json<T>, E>
120    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
149/// Find a `State<S>` parameter and return the inner `S`.
150fn 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
163/// Find the first `Json<T>` parameter and return the inner `T`.
164fn 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// ---------------------------------------------------------------------------
176// Tests
177// ---------------------------------------------------------------------------
178
179#[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    // --- valid signatures ---
194
195    #[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        // output_type is unit () for SSE handlers
286        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    // --- invalid signatures ---
299
300    #[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}