Support using a different rejection for #[derive(FromRequest)] (#1256)

This commit is contained in:
David Pedersen
2022-08-12 16:05:27 +00:00
committed by GitHub
parent a8e80bcb97
commit ac7037d282
18 changed files with 750 additions and 45 deletions
+213 -34
View File
@@ -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)
}
}
})