mirror of
https://github.com/tokio-rs/axum.git
synced 2026-09-08 00:00:24 +02:00
Support using a different rejection for #[derive(FromRequest)] (#1256)
This commit is contained in:
+213
-34
@@ -5,7 +5,7 @@ use self::attr::{
|
||||
use heck::ToUpperCamelCase;
|
||||
use proc_macro2::{Span, TokenStream};
|
||||
use quote::{format_ident, quote, quote_spanned};
|
||||
use syn::{punctuated::Punctuated, spanned::Spanned, Token};
|
||||
use syn::{punctuated::Punctuated, spanned::Spanned, Ident, Token};
|
||||
|
||||
mod attr;
|
||||
|
||||
@@ -22,21 +22,45 @@ pub(crate) fn expand(item: syn::Item) -> syn::Result<TokenStream> {
|
||||
struct_token: _,
|
||||
} = item;
|
||||
|
||||
error_on_generics(generics)?;
|
||||
let generic_ident = parse_single_generic_type_on_struct(generics, &fields)?;
|
||||
|
||||
match parse_container_attrs(&attrs)? {
|
||||
FromRequestContainerAttr::Via(path) => {
|
||||
impl_struct_by_extracting_all_at_once(ident, fields, path)
|
||||
FromRequestContainerAttr::Via { path, rejection } => {
|
||||
impl_struct_by_extracting_all_at_once(
|
||||
ident,
|
||||
fields,
|
||||
path,
|
||||
rejection,
|
||||
generic_ident,
|
||||
)
|
||||
}
|
||||
FromRequestContainerAttr::RejectionDerive(_, opt_outs) => {
|
||||
impl_struct_by_extracting_each_field(ident, fields, vis, opt_outs)
|
||||
error_on_generic_ident(generic_ident)?;
|
||||
|
||||
impl_struct_by_extracting_each_field(ident, fields, vis, opt_outs, None)
|
||||
}
|
||||
FromRequestContainerAttr::Rejection(rejection) => {
|
||||
error_on_generic_ident(generic_ident)?;
|
||||
|
||||
impl_struct_by_extracting_each_field(
|
||||
ident,
|
||||
fields,
|
||||
vis,
|
||||
RejectionDeriveOptOuts::default(),
|
||||
Some(rejection),
|
||||
)
|
||||
}
|
||||
FromRequestContainerAttr::None => {
|
||||
error_on_generic_ident(generic_ident)?;
|
||||
|
||||
impl_struct_by_extracting_each_field(
|
||||
ident,
|
||||
fields,
|
||||
vis,
|
||||
RejectionDeriveOptOuts::default(),
|
||||
None,
|
||||
)
|
||||
}
|
||||
FromRequestContainerAttr::None => impl_struct_by_extracting_each_field(
|
||||
ident,
|
||||
fields,
|
||||
vis,
|
||||
RejectionDeriveOptOuts::default(),
|
||||
),
|
||||
}
|
||||
}
|
||||
syn::Item::Enum(item) => {
|
||||
@@ -50,11 +74,19 @@ pub(crate) fn expand(item: syn::Item) -> syn::Result<TokenStream> {
|
||||
variants,
|
||||
} = item;
|
||||
|
||||
error_on_generics(generics)?;
|
||||
const GENERICS_ERROR: &str = "`#[derive(FromRequest)] on enums don't support generics";
|
||||
|
||||
if !generics.params.is_empty() {
|
||||
return Err(syn::Error::new_spanned(generics, GENERICS_ERROR));
|
||||
}
|
||||
|
||||
if let Some(where_clause) = generics.where_clause {
|
||||
return Err(syn::Error::new_spanned(where_clause, GENERICS_ERROR));
|
||||
}
|
||||
|
||||
match parse_container_attrs(&attrs)? {
|
||||
FromRequestContainerAttr::Via(path) => {
|
||||
impl_enum_by_extracting_all_at_once(ident, variants, path)
|
||||
FromRequestContainerAttr::Via { path, rejection } => {
|
||||
impl_enum_by_extracting_all_at_once(ident, variants, path, rejection)
|
||||
}
|
||||
FromRequestContainerAttr::RejectionDerive(rejection_derive, _) => {
|
||||
Err(syn::Error::new_spanned(
|
||||
@@ -62,6 +94,10 @@ pub(crate) fn expand(item: syn::Item) -> syn::Result<TokenStream> {
|
||||
"cannot use `rejection_derive` on enums",
|
||||
))
|
||||
}
|
||||
FromRequestContainerAttr::Rejection(rejection) => Err(syn::Error::new_spanned(
|
||||
rejection,
|
||||
"cannot use `rejection` without `via`",
|
||||
)),
|
||||
FromRequestContainerAttr::None => Err(syn::Error::new(
|
||||
Span::call_site(),
|
||||
"missing `#[from_request(via(...))]`",
|
||||
@@ -72,18 +108,90 @@ pub(crate) fn expand(item: syn::Item) -> syn::Result<TokenStream> {
|
||||
}
|
||||
}
|
||||
|
||||
fn error_on_generics(generics: syn::Generics) -> syn::Result<()> {
|
||||
const GENERICS_ERROR: &str = "`#[derive(FromRequest)] doesn't support generics";
|
||||
|
||||
if !generics.params.is_empty() {
|
||||
return Err(syn::Error::new_spanned(generics, GENERICS_ERROR));
|
||||
}
|
||||
|
||||
fn parse_single_generic_type_on_struct(
|
||||
generics: syn::Generics,
|
||||
fields: &syn::Fields,
|
||||
) -> syn::Result<Option<Ident>> {
|
||||
if let Some(where_clause) = generics.where_clause {
|
||||
return Err(syn::Error::new_spanned(where_clause, GENERICS_ERROR));
|
||||
return Err(syn::Error::new_spanned(
|
||||
where_clause,
|
||||
"#[derive(FromRequest)] doesn't support structs with `where` clauses",
|
||||
));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
match generics.params.len() {
|
||||
0 => Ok(None),
|
||||
1 => {
|
||||
let param = generics.params.first().unwrap();
|
||||
let ty_ident = match param {
|
||||
syn::GenericParam::Type(ty) => &ty.ident,
|
||||
syn::GenericParam::Lifetime(lifetime) => {
|
||||
return Err(syn::Error::new_spanned(
|
||||
lifetime,
|
||||
"#[derive(FromRequest)] doesn't support structs that are generic over lifetimes",
|
||||
));
|
||||
}
|
||||
syn::GenericParam::Const(konst) => {
|
||||
return Err(syn::Error::new_spanned(
|
||||
konst,
|
||||
"#[derive(FromRequest)] doesn't support structs that have const generics",
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
match fields {
|
||||
syn::Fields::Named(fields_named) => {
|
||||
return Err(syn::Error::new_spanned(
|
||||
fields_named,
|
||||
"#[derive(FromRequest)] doesn't support named fields for generic structs. Use a tuple struct instead",
|
||||
));
|
||||
}
|
||||
syn::Fields::Unnamed(fields_unnamed) => {
|
||||
if fields_unnamed.unnamed.len() != 1 {
|
||||
return Err(syn::Error::new_spanned(
|
||||
fields_unnamed,
|
||||
"#[derive(FromRequest)] only supports generics on tuple structs that have exactly one field",
|
||||
));
|
||||
}
|
||||
|
||||
let field = fields_unnamed.unnamed.first().unwrap();
|
||||
|
||||
if let syn::Type::Path(type_path) = &field.ty {
|
||||
if type_path
|
||||
.path
|
||||
.get_ident()
|
||||
.map_or(true, |field_type_ident| field_type_ident != ty_ident)
|
||||
{
|
||||
return Err(syn::Error::new_spanned(
|
||||
type_path,
|
||||
"#[derive(FromRequest)] only supports generics on tuple structs that have exactly one field of the generic type",
|
||||
));
|
||||
}
|
||||
} else {
|
||||
return Err(syn::Error::new_spanned(&field.ty, "Expected type path"));
|
||||
}
|
||||
}
|
||||
syn::Fields::Unit => return Ok(None),
|
||||
}
|
||||
|
||||
Ok(Some(ty_ident.clone()))
|
||||
}
|
||||
_ => Err(syn::Error::new_spanned(
|
||||
generics,
|
||||
"#[derive(FromRequest)] only supports 0 or 1 generic type parameters",
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn error_on_generic_ident(generic_ident: Option<Ident>) -> syn::Result<()> {
|
||||
if let Some(generic_ident) = generic_ident {
|
||||
Err(syn::Error::new_spanned(
|
||||
generic_ident,
|
||||
"#[derive(FromRequest)] only supports generics when used with #[from_request(via)]",
|
||||
))
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn impl_struct_by_extracting_each_field(
|
||||
@@ -91,10 +199,14 @@ fn impl_struct_by_extracting_each_field(
|
||||
fields: syn::Fields,
|
||||
vis: syn::Visibility,
|
||||
rejection_derive_opt_outs: RejectionDeriveOptOuts,
|
||||
rejection: Option<syn::Path>,
|
||||
) -> syn::Result<TokenStream> {
|
||||
let extract_fields = extract_fields(&fields)?;
|
||||
let extract_fields = extract_fields(&fields, &rejection)?;
|
||||
|
||||
let (rejection_ident, rejection) = if has_no_fields(&fields) {
|
||||
let (rejection_ident, rejection) = if let Some(rejection) = rejection {
|
||||
let rejection_ident = syn::parse_quote!(#rejection);
|
||||
(rejection_ident, None)
|
||||
} else if has_no_fields(&fields) {
|
||||
(syn::parse_quote!(::std::convert::Infallible), None)
|
||||
} else {
|
||||
let rejection_ident = rejection_ident(&ident);
|
||||
@@ -140,7 +252,10 @@ fn rejection_ident(ident: &syn::Ident) -> syn::Type {
|
||||
syn::parse_quote!(#ident)
|
||||
}
|
||||
|
||||
fn extract_fields(fields: &syn::Fields) -> syn::Result<Vec<TokenStream>> {
|
||||
fn extract_fields(
|
||||
fields: &syn::Fields,
|
||||
rejection: &Option<syn::Path>,
|
||||
) -> syn::Result<Vec<TokenStream>> {
|
||||
fields
|
||||
.iter()
|
||||
.enumerate()
|
||||
@@ -190,12 +305,18 @@ fn extract_fields(fields: &syn::Fields) -> syn::Result<Vec<TokenStream>> {
|
||||
},
|
||||
})
|
||||
} else {
|
||||
let map_err = if let Some(rejection) = rejection {
|
||||
quote! { <#rejection as ::std::convert::From<_>>::from }
|
||||
} else {
|
||||
quote! { Self::Rejection::#rejection_variant_name }
|
||||
};
|
||||
|
||||
Ok(quote_spanned! {ty_span=>
|
||||
#member: {
|
||||
::axum::extract::FromRequest::from_request(req)
|
||||
.await
|
||||
.map(#into_inner)
|
||||
.map_err(Self::Rejection::#rejection_variant_name)?
|
||||
.map_err(#map_err)?
|
||||
},
|
||||
})
|
||||
}
|
||||
@@ -462,6 +583,8 @@ fn impl_struct_by_extracting_all_at_once(
|
||||
ident: syn::Ident,
|
||||
fields: syn::Fields,
|
||||
path: syn::Path,
|
||||
rejection: Option<syn::Path>,
|
||||
generic_ident: Option<Ident>,
|
||||
) -> syn::Result<TokenStream> {
|
||||
let fields = match fields {
|
||||
syn::Fields::Named(fields) => fields.named.into_iter(),
|
||||
@@ -482,23 +605,69 @@ fn impl_struct_by_extracting_all_at_once(
|
||||
|
||||
let path_span = path.span();
|
||||
|
||||
let associated_rejection_type = if let Some(rejection) = &rejection {
|
||||
quote! { #rejection }
|
||||
} else {
|
||||
quote! {
|
||||
<#path<Self> as ::axum::extract::FromRequest<B>>::Rejection
|
||||
}
|
||||
};
|
||||
|
||||
let rejection_bound = rejection.as_ref().map(|rejection| {
|
||||
if generic_ident.is_some() {
|
||||
quote! {
|
||||
#rejection: ::std::convert::From<<#path<T> as ::axum::extract::FromRequest<B>>::Rejection>,
|
||||
}
|
||||
} else {
|
||||
quote! {
|
||||
#rejection: ::std::convert::From<<#path<Self> as ::axum::extract::FromRequest<B>>::Rejection>,
|
||||
}
|
||||
}
|
||||
}).unwrap_or_default();
|
||||
|
||||
let impl_generics = if generic_ident.is_some() {
|
||||
quote! { B, T }
|
||||
} else {
|
||||
quote! { B }
|
||||
};
|
||||
|
||||
let type_generics = generic_ident
|
||||
.is_some()
|
||||
.then(|| quote! { <T> })
|
||||
.unwrap_or_default();
|
||||
|
||||
let via_type_generics = if generic_ident.is_some() {
|
||||
quote! { T }
|
||||
} else {
|
||||
quote! { Self }
|
||||
};
|
||||
|
||||
let value_to_self = if generic_ident.is_some() {
|
||||
quote! {
|
||||
#ident(value)
|
||||
}
|
||||
} else {
|
||||
quote! { value }
|
||||
};
|
||||
|
||||
Ok(quote_spanned! {path_span=>
|
||||
#[::axum::async_trait]
|
||||
#[automatically_derived]
|
||||
impl<B> ::axum::extract::FromRequest<B> for #ident
|
||||
impl<#impl_generics> ::axum::extract::FromRequest<B> for #ident #type_generics
|
||||
where
|
||||
B: ::axum::body::HttpBody + ::std::marker::Send + 'static,
|
||||
B::Data: ::std::marker::Send,
|
||||
B::Error: ::std::convert::Into<::axum::BoxError>,
|
||||
#path<#via_type_generics>: ::axum::extract::FromRequest<B>,
|
||||
#rejection_bound
|
||||
B: ::std::marker::Send,
|
||||
{
|
||||
type Rejection = <#path<Self> as ::axum::extract::FromRequest<B>>::Rejection;
|
||||
type Rejection = #associated_rejection_type;
|
||||
|
||||
async fn from_request(
|
||||
req: &mut ::axum::extract::RequestParts<B>,
|
||||
) -> ::std::result::Result<Self, Self::Rejection> {
|
||||
::axum::extract::FromRequest::<B>::from_request(req)
|
||||
.await
|
||||
.map(|#path(inner)| inner)
|
||||
.map(|#path(value)| #value_to_self)
|
||||
.map_err(::std::convert::From::from)
|
||||
}
|
||||
}
|
||||
})
|
||||
@@ -508,6 +677,7 @@ fn impl_enum_by_extracting_all_at_once(
|
||||
ident: syn::Ident,
|
||||
variants: Punctuated<syn::Variant, Token![,]>,
|
||||
path: syn::Path,
|
||||
rejection: Option<syn::Path>,
|
||||
) -> syn::Result<TokenStream> {
|
||||
for variant in variants {
|
||||
let FromRequestFieldAttr { via } = parse_field_attrs(&variant.attrs)?;
|
||||
@@ -535,6 +705,14 @@ fn impl_enum_by_extracting_all_at_once(
|
||||
}
|
||||
}
|
||||
|
||||
let associated_rejection_type = if let Some(rejection) = rejection {
|
||||
quote! { #rejection }
|
||||
} else {
|
||||
quote! {
|
||||
<#path<Self> as ::axum::extract::FromRequest<B>>::Rejection
|
||||
}
|
||||
};
|
||||
|
||||
let path_span = path.span();
|
||||
|
||||
Ok(quote_spanned! {path_span=>
|
||||
@@ -546,7 +724,7 @@ fn impl_enum_by_extracting_all_at_once(
|
||||
B::Data: ::std::marker::Send,
|
||||
B::Error: ::std::convert::Into<::axum::BoxError>,
|
||||
{
|
||||
type Rejection = <#path<Self> as ::axum::extract::FromRequest<B>>::Rejection;
|
||||
type Rejection = #associated_rejection_type;
|
||||
|
||||
async fn from_request(
|
||||
req: &mut ::axum::extract::RequestParts<B>,
|
||||
@@ -554,6 +732,7 @@ fn impl_enum_by_extracting_all_at_once(
|
||||
::axum::extract::FromRequest::<B>::from_request(req)
|
||||
.await
|
||||
.map(|#path(inner)| inner)
|
||||
.map_err(::std::convert::From::from)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user