use proc_macro2::Span;
use syn::{FnArg, ItemFn, ReturnType, Type, spanned::Spanned};
use crate::{
errors::{Error, Result},
types::{JSON, RESULT, SSE, STATE, try_extract_wrapper},
};
#[derive(Debug)]
pub struct HandlerSignature {
pub fn_name: String,
pub fn_span: Span,
pub state_type: Option<Type>,
pub input_type: Option<Type>,
pub output_type: Type,
pub error_type: Option<Type>,
pub is_streaming: bool,
pub is_async: bool,
}
pub fn extract_handler_signature(func: &ItemFn) -> Result<HandlerSignature> {
let fn_name = func.sig.ident.to_string();
let fn_span = func.sig.ident.span();
let is_async = func.sig.asyncness.is_some();
let (output_type, error_type, is_streaming) = extract_return_types(&func.sig.output, &fn_name)?;
let state_type = extract_state_param(&func.sig.inputs);
let input_type = extract_json_param(&func.sig.inputs);
Ok(HandlerSignature {
fn_name,
fn_span,
state_type,
input_type,
output_type,
error_type,
is_streaming,
is_async,
})
}
fn extract_return_types(
return_type: &ReturnType,
fn_name: &str,
) -> Result<(Type, Option<Type>, bool)> {
let ty = match return_type {
ReturnType::Default => {
return Err(Error::missing_return_type(
proc_macro2::Span::call_site(),
fn_name,
));
}
ReturnType::Type(_, ty) => ty.as_ref(),
};
if try_extract_wrapper(ty, SSE).is_some() {
let unit: Type = syn::parse_quote! { () };
return Ok((unit, None, true));
}
if let Some(m) = try_extract_wrapper(ty, JSON) {
let output = m
.first_type()
.ok_or_else(|| Error::empty_generic_args(ty.span(), JSON))?
.clone();
return Ok((output, None, false));
}
if let Some(result_match) = try_extract_wrapper(ty, RESULT) {
let first = result_match
.first_type()
.ok_or_else(|| Error::empty_generic_args(ty.span(), RESULT))?;
let json_match = try_extract_wrapper(first, JSON).ok_or_else(|| {
Error::invalid_handler_sig(
first.span(),
fn_name,
"Result's first type argument must be Json<T>",
)
})?;
let output = json_match
.first_type()
.ok_or_else(|| Error::empty_generic_args(first.span(), JSON))?
.clone();
let error_type = result_match.second_type().cloned();
return Ok((output, error_type, false));
}
Err(Error::invalid_handler_sig(
ty.span(),
fn_name,
"return type must be Json<T>, Result<Json<T>, E>, or Sse<impl Stream<...>>",
))
}
fn extract_state_param(
inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>,
) -> Option<Type> {
for arg in inputs {
if let FnArg::Typed(pat_type) = arg
&& let Some(m) = try_extract_wrapper(&pat_type.ty, STATE)
{
return m.first_type().cloned();
}
}
None
}
fn extract_json_param(inputs: &syn::punctuated::Punctuated<FnArg, syn::Token![,]>) -> Option<Type> {
for arg in inputs {
if let FnArg::Typed(pat_type) = arg
&& let Some(m) = try_extract_wrapper(&pat_type.ty, JSON)
{
return m.first_type().cloned();
}
}
None
}
#[cfg(test)]
mod tests {
use super::*;
use crate::errors::type_display;
use syn::parse_quote;
fn sig(func: ItemFn) -> HandlerSignature {
extract_handler_signature(&func).unwrap()
}
fn sig_err(func: ItemFn) -> Error {
extract_handler_signature(&func).unwrap_err()
}
#[test]
fn json_return_only() {
let f: ItemFn = parse_quote! {
async fn handler() -> Json<Planet> {}
};
let s = sig(f);
assert_eq!(s.fn_name, "handler");
assert_eq!(type_display(&s.output_type), "Planet");
assert!(s.error_type.is_none());
assert!(s.input_type.is_none());
assert!(s.state_type.is_none());
assert!(s.is_async);
assert!(!s.is_streaming);
}
#[test]
fn result_json_return() {
let f: ItemFn = parse_quote! {
async fn handler() -> Result<Json<Planet>, AppError> {}
};
let s = sig(f);
assert_eq!(type_display(&s.output_type), "Planet");
assert_eq!(type_display(s.error_type.as_ref().unwrap()), "AppError");
}
#[test]
fn qualified_result_return() {
let f: ItemFn = parse_quote! {
async fn handler() -> std::result::Result<Json<Planet>, AppError> {}
};
let s = sig(f);
assert_eq!(type_display(&s.output_type), "Planet");
}
#[test]
fn qualified_json_return() {
let f: ItemFn = parse_quote! {
async fn handler() -> axum::extract::Json<Planet> {}
};
let s = sig(f);
assert_eq!(type_display(&s.output_type), "Planet");
}
#[test]
fn state_param_extracted() {
let f: ItemFn = parse_quote! {
async fn handler(State(db): State<Db>) -> Json<Planet> {}
};
let s = sig(f);
assert_eq!(type_display(s.state_type.as_ref().unwrap()), "Db");
}
#[test]
fn json_param_extracted() {
let f: ItemFn = parse_quote! {
async fn handler(Json(body): Json<CreatePlanet>) -> Json<Planet> {}
};
let s = sig(f);
assert_eq!(type_display(s.input_type.as_ref().unwrap()), "CreatePlanet");
}
#[test]
fn both_state_and_json_params() {
let f: ItemFn = parse_quote! {
async fn handler(State(db): State<Db>, Json(body): Json<CreatePlanet>) -> Result<Json<Planet>, AppError> {}
};
let s = sig(f);
assert_eq!(type_display(s.state_type.as_ref().unwrap()), "Db");
assert_eq!(type_display(s.input_type.as_ref().unwrap()), "CreatePlanet");
assert_eq!(type_display(&s.output_type), "Planet");
assert!(s.error_type.is_some());
}
#[test]
fn sync_function_allowed() {
let f: ItemFn = parse_quote! {
fn handler() -> Json<Planet> {}
};
let s = sig(f);
assert!(!s.is_async);
}
#[test]
fn sse_return_is_streaming() {
let f: ItemFn = parse_quote! {
async fn stream_events(State(_state): State<AppState>) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {}
};
let s = sig(f);
assert!(s.is_streaming);
assert!(s.error_type.is_none());
assert_eq!(type_display(&s.output_type), "()");
}
#[test]
fn qualified_sse_return_is_streaming() {
let f: ItemFn = parse_quote! {
async fn handler() -> axum::response::Sse<SomeStream> {}
};
let s = sig(f);
assert!(s.is_streaming);
}
#[test]
fn no_return_type_error() {
let f: ItemFn = parse_quote! {
async fn handler() {}
};
let err = sig_err(f);
assert!(err.to_string().contains("has no return type"));
}
#[test]
fn bare_type_return_error() {
let f: ItemFn = parse_quote! {
async fn handler() -> Vec<Planet> {}
};
let err = sig_err(f);
assert!(err.to_string().contains("return type must be Json<T>"));
}
#[test]
fn result_without_json_inner_error() {
let f: ItemFn = parse_quote! {
async fn handler() -> Result<Vec<Planet>, AppError> {}
};
let err = sig_err(f);
assert!(
err.to_string()
.contains("Result's first type argument must be Json<T>")
);
}
}