use crate::{controllers::shared::CMetaStack, interceptor::InterceptorArgs};
use proc_macro::TokenStream;
use proc_macro2::{Span, TokenStream as TokenStream2};
use quote::quote;
use syn::spanned::Spanned;
use syn::{Attribute, GenericArgument, ItemFn, LitStr, PathArguments, ReturnType, Type};
#[derive(Clone, Copy, PartialEq, Eq)]
pub enum RequestMode {
None,
Buffered,
Streaming,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum ReturnKind {
Passthrough,
Serialize,
Empty,
}
#[derive(Clone, Copy)]
pub enum HttpMethod {
Get,
Post,
Put,
Delete,
Patch,
Head,
Options,
Trace,
Connect,
}
impl HttpMethod {
pub fn from_attr_name(method: &str) -> syn::Result<Self> {
match method {
"GET" => Ok(Self::Get),
"POST" => Ok(Self::Post),
"PUT" => Ok(Self::Put),
"DELETE" => Ok(Self::Delete),
"PATCH" => Ok(Self::Patch),
"HEAD" => Ok(Self::Head),
"OPTIONS" => Ok(Self::Options),
"TRACE" => Ok(Self::Trace),
"CONNECT" => Ok(Self::Connect),
_ => Err(syn::Error::new(
Span::call_site(),
format!("Unsupported HTTP method `{method}`"),
)),
}
}
pub fn as_str(self) -> &'static str {
match self {
Self::Get => "GET",
Self::Post => "POST",
Self::Put => "PUT",
Self::Delete => "DELETE",
Self::Patch => "PATCH",
Self::Head => "HEAD",
Self::Options => "OPTIONS",
Self::Trace => "TRACE",
Self::Connect => "CONNECT",
}
}
pub fn routing_fn_tokens(self) -> TokenStream2 {
match self {
Self::Get => quote! { get },
Self::Post => quote! { post },
Self::Put => quote! { put },
Self::Delete => quote! { delete },
Self::Patch => quote! { patch },
Self::Head => quote! { head },
Self::Options => quote! { options },
Self::Trace => quote! { trace },
Self::Connect => quote! { connect },
}
}
pub fn default_status_code(self) -> u16 {
match self {
Self::Post => 201,
_ => 200,
}
}
}
pub struct WebRouteContext {
pub controller_name: String,
pub controller_interceptors: Vec<InterceptorArgs>,
}
impl WebRouteContext {
fn from_cmeta(method: &str) -> syn::Result<Self> {
let Some(controller_name) = CMetaStack::get("web", "controller_name") else {
let error = format!(
"\n[ERROR] The #[{}] attribute must be used inside a #[controller] impl block.\n\
\n\
Make sure:\n\
1. The struct has #[controller(kind = Controller::Web, path = \"/path\")] attribute\n\
2. The struct is defined BEFORE the impl block\n\
3. The impl block is for the same struct\n",
method
);
return Err(syn::Error::new(Span::call_site(), error));
};
let controller_interceptors = CMetaStack::get_list("web", "controller_interceptors")
.unwrap_or_default()
.into_iter()
.map(|interceptor| syn::parse_str::<InterceptorArgs>(&interceptor))
.collect::<syn::Result<Vec<_>>>()?;
Ok(Self {
controller_name,
controller_interceptors,
})
}
}
pub struct ParsedRouteAttribute {
pub method: HttpMethod,
pub path: String,
pub function: ItemFn,
pub interceptors: Vec<InterceptorArgs>,
pub request_mode: RequestMode,
pub context: WebRouteContext,
pub is_result_return: bool,
pub return_kind: ReturnKind,
pub status_code: u16,
}
impl ParsedRouteAttribute {
pub fn parse(method: &str, attr: TokenStream, item: TokenStream) -> syn::Result<Self> {
let method = HttpMethod::from_attr_name(method)?;
let path = Self::parse_path(attr)?;
let mut input_fn = Self::parse_function(item)?;
let request_mode = Self::infer_request_mode(&input_fn)?;
let (interceptors, retained_attrs) = Self::extract_interceptors(&input_fn)?;
input_fn.attrs = retained_attrs;
let context = Self::resolve_context(method.as_str())?;
let (is_result_return, return_kind, status_code) =
Self::extract_result_info(&input_fn.sig, &method);
Ok(Self {
method,
path,
function: input_fn,
interceptors,
request_mode,
context,
is_result_return,
return_kind,
status_code,
})
}
fn parse_path(attr: TokenStream) -> syn::Result<String> {
if attr.is_empty() {
return Err(syn::Error::new(
Span::call_site(),
"Route path is required, e.g., #[get(\"/path\")]. If you want the root path, use \"/\".",
));
}
Ok(syn::parse::<LitStr>(attr)?.value())
}
fn parse_function(item: TokenStream) -> syn::Result<ItemFn> {
syn::parse::<ItemFn>(item)
}
fn extract_interceptors(
input_fn: &ItemFn,
) -> syn::Result<(Vec<InterceptorArgs>, Vec<Attribute>)> {
let mut interceptors = Vec::new();
let mut retained_attrs = Vec::new();
for attr in input_fn.attrs.iter() {
if attr.path().is_ident("interceptor") {
interceptors.push(attr.parse_args::<InterceptorArgs>()?);
} else {
retained_attrs.push(attr.clone());
}
}
Ok((interceptors, retained_attrs))
}
fn resolve_context(method: &str) -> syn::Result<WebRouteContext> {
WebRouteContext::from_cmeta(method)
}
fn infer_request_mode(input_fn: &ItemFn) -> syn::Result<RequestMode> {
let mut mode = RequestMode::None;
for arg in &input_fn.sig.inputs {
let syn::FnArg::Typed(pat_type) = arg else {
continue;
};
let arg_mode = Self::request_mode_from_type(&pat_type.ty);
if arg_mode == RequestMode::None {
continue;
}
if mode != RequestMode::None && mode != arg_mode {
return Err(syn::Error::new(
pat_type.ty.span(),
"A route handler cannot use both `Request` and `StreamRequest` in the same signature",
));
}
mode = arg_mode;
}
Ok(mode)
}
fn request_mode_from_type(ty: &Type) -> RequestMode {
let Type::Path(type_path) = ty else {
return RequestMode::None;
};
let Some(last_segment) = type_path.path.segments.last() else {
return RequestMode::None;
};
if last_segment.ident == "Request" {
return RequestMode::Buffered;
}
if last_segment.ident == "StreamRequest" {
return RequestMode::Streaming;
}
RequestMode::None
}
fn extract_result_info(sig: &syn::Signature, method: &HttpMethod) -> (bool, ReturnKind, u16) {
let status_code = method.default_status_code();
let return_type = match &sig.output {
ReturnType::Type(_, ty) => ty.as_ref(),
_ => return (false, ReturnKind::Passthrough, status_code),
};
let type_path = match return_type {
Type::Path(type_path) => type_path,
_ => return (false, ReturnKind::Passthrough, status_code),
};
let last_segment = match type_path.path.segments.last() {
Some(seg) => seg,
None => return (false, ReturnKind::Passthrough, status_code),
};
let ident = last_segment.ident.to_string();
let ok_type = match ident.as_str() {
"Result" => {
let args = match &last_segment.arguments {
PathArguments::AngleBracketed(args) => &args.args,
_ => return (false, ReturnKind::Passthrough, status_code),
};
if args.len() < 2 {
return (false, ReturnKind::Passthrough, status_code);
}
match &args[0] {
GenericArgument::Type(ty) => ty.clone(),
_ => return (false, ReturnKind::Passthrough, status_code),
}
}
"WebResult" => {
match &last_segment.arguments {
PathArguments::AngleBracketed(args) => {
if args.args.is_empty() {
syn::parse_quote!(JsonResponse)
} else {
match &args.args[0] {
GenericArgument::Type(ty) => ty.clone(),
_ => return (false, ReturnKind::Passthrough, status_code),
}
}
}
_ => {
syn::parse_quote!(JsonResponse)
}
}
}
_ => return (false, ReturnKind::Passthrough, status_code),
};
let return_kind = Self::classify_return_type(&ok_type);
(true, return_kind, status_code)
}
fn classify_return_type(ty: &Type) -> ReturnKind {
match ty {
Type::Tuple(tuple) if tuple.elems.is_empty() => ReturnKind::Empty,
Type::Path(type_path) => {
let Some(last_segment) = type_path.path.segments.last() else {
return ReturnKind::Serialize;
};
match last_segment.ident.to_string().as_str() {
"JsonResponse" | "File" | "Redirect" => ReturnKind::Passthrough,
_ => ReturnKind::Serialize,
}
}
_ => ReturnKind::Serialize,
}
}
}