diff --git a/api-error-derive/src/expand.rs b/api-error-derive/src/expand.rs index eeab368..3fe4467 100644 --- a/api-error-derive/src/expand.rs +++ b/api-error-derive/src/expand.rs @@ -10,6 +10,7 @@ use crate::{VariantAttr, parser}; struct Expansion { status_code: TokenStream, message: TokenStream, + extended: TokenStream, } pub fn expand(input: DeriveInput) -> TokenStream { @@ -25,6 +26,7 @@ pub fn expand(input: DeriveInput) -> TokenStream { let Expansion { status_code, message, + extended, } = match tokens { Ok(ts) => ts, Err(err) => return err.to_compile_error(), @@ -33,6 +35,19 @@ pub fn expand(input: DeriveInput) -> TokenStream { let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); let ident = &input.ident; + #[cfg(feature = "axum")] + let extended_impl = Some(quote! { + fn extended(&self) -> ::std::option::Option<::api_error::__serde_json::Value> { + #extended + } + }); + + #[cfg(not(feature = "axum"))] + let extended_impl: Option = { + let _ = extended; + None + }; + let api_err_impl = quote! { #[automatically_derived] impl #impl_generics ApiError for #ident #ty_generics #where_clause { @@ -43,6 +58,8 @@ pub fn expand(input: DeriveInput) -> TokenStream { fn message<'a>(&'a self) -> ::std::borrow::Cow<'a, str> { #message } + + #extended_impl } }; @@ -93,6 +110,13 @@ fn expand_struct(ident: &Ident, data: DataStruct, attrs: &[Attribute]) -> syn::R } }; + let extended = match &attr { + VariantAttr::Transparent => quote! { ApiError::extended(&self.0) }, + VariantAttr::InheritMsg { .. } | VariantAttr::Custom { .. } => { + quote! { ::std::option::Option::None } + } + }; + let message = match (attr, data.fields) { (VariantAttr::Transparent, fields) if fields.len() == 1 => { quote! { ApiError::message(&self.0) } @@ -145,6 +169,7 @@ fn expand_struct(ident: &Ident, data: DataStruct, attrs: &[Attribute]) -> syn::R Ok(Expansion { status_code, message, + extended, }) } @@ -155,14 +180,20 @@ fn expand_enum(data: DataEnum) -> syn::Result { .map(expand_enum_variant) .collect::>>()?; - let (status_arms, message_arms): (Vec<_>, Vec<_>) = variant_expansions - .into_iter() - .map(|v| (v.status_code, v.message)) - .unzip(); + let mut status_arms = Vec::with_capacity(variant_expansions.len()); + let mut message_arms = Vec::with_capacity(variant_expansions.len()); + let mut extended_arms = Vec::with_capacity(variant_expansions.len()); + + for v in variant_expansions { + status_arms.push(v.status_code); + message_arms.push(v.message); + extended_arms.push(v.extended); + } Ok(Expansion { status_code: quote! { match self { #(#status_arms),* } }, message: quote! { match self { #(#message_arms),* } }, + extended: quote! { match self { #(#extended_arms),* } }, }) } @@ -171,11 +202,13 @@ fn expand_enum_variant(v: Variant) -> syn::Result { let attr_args = parser::parse_variant_attrs(&v.attrs)?; let status_arm = expand_status_arm(&v.ident, &v.fields, &attr_args)?; + let extended_arm = expand_extended_arm(&v.ident, &v.fields, &attr_args)?; let message_arm = expand_message_arm(&v.ident, &v.fields, attr_args)?; Ok(Expansion { status_code: status_arm, message: message_arm, + extended: extended_arm, }) } @@ -275,3 +308,30 @@ fn expand_message_arm( Ok(quote! { #message_pat => #message_arm }) } + +/// Generate a match statement for the `extended` method. +fn expand_extended_arm( + variant_ident: &Ident, + fields: &Fields, + attr: &VariantAttr, +) -> syn::Result { + let extended_pat = expand_variant_pattern(variant_ident, fields); + + let extended_arm = match (fields, attr) { + // transparent expansion, forward to the inner field + (Fields::Unnamed(fields), VariantAttr::Transparent) if fields.unnamed.len() == 1 => { + quote! { ApiError::extended(__field0) } + } + (_, VariantAttr::Transparent) => Err(syn::Error::new_spanned( + variant_ident, + "the `#[api_error(transparent)]` attribute is only allowed on unamed variants with only one field", + ))?, + + // non-transparent variants keep the trait default + (_, VariantAttr::Custom { .. } | VariantAttr::InheritMsg { .. }) => { + quote! { ::std::option::Option::None } + } + }; + + Ok(quote! { #extended_pat => #extended_arm }) +} diff --git a/api-error/src/lib.rs b/api-error/src/lib.rs index 2f23b75..580f1b5 100644 --- a/api-error/src/lib.rs +++ b/api-error/src/lib.rs @@ -232,6 +232,10 @@ use http::StatusCode; #[doc(hidden)] pub use ::http as __http; +#[cfg(feature = "axum")] +#[doc(hidden)] +pub use ::serde_json as __serde_json; + /// Derive macro for implementing [`ApiError`] on enums and structs. /// /// This macro generates an [`ApiError`] implementation based on diff --git a/api-error/tests/responder.rs b/api-error/tests/responder.rs index e3424a2..b6bd9aa 100644 --- a/api-error/tests/responder.rs +++ b/api-error/tests/responder.rs @@ -84,3 +84,53 @@ fn extended_forwarded_through_reference() { let by_ref: &ValidationError = &err; assert_eq!(by_ref.extended(), Some(json!({ "field": "email" }))); } + +// A derived transparent wrapper must forward `extended` to the inner error, +// just like it forwards `status_code` and `message`. +#[derive(Debug, thiserror::Error, ApiError)] +#[error(transparent)] +#[api_error(transparent)] +struct TransparentStruct(ValidationError); + +#[derive(Debug, thiserror::Error, ApiError)] +enum TransparentEnum { + #[error(transparent)] + #[api_error(transparent)] + Validation(ValidationError), + + #[error("plain")] + #[api_error(status_code = 400, message = "plain")] + Plain, +} + +#[test] +fn transparent_struct_forwards_extended() { + let err = TransparentStruct(ValidationError); + assert_eq!(err.status_code(), StatusCode::UNPROCESSABLE_ENTITY); + assert_eq!(err.message().as_ref(), "validation failed"); + assert_eq!(err.extended(), Some(json!({ "field": "email" }))); + + let body = serde_json::to_string(&ApiErrorResponse::new(&err)).unwrap(); + assert_eq!( + body, + r#"{"message":"validation failed","extended":{"field":"email"}}"# + ); +} + +#[test] +fn transparent_enum_forwards_extended() { + let err = TransparentEnum::Validation(ValidationError); + assert_eq!(err.extended(), Some(json!({ "field": "email" }))); + + let body = serde_json::to_string(&ApiErrorResponse::new(&err)).unwrap(); + assert_eq!( + body, + r#"{"message":"validation failed","extended":{"field":"email"}}"# + ); +} + +#[test] +fn non_transparent_enum_variant_keeps_extended_none() { + let err = TransparentEnum::Plain; + assert!(err.extended().is_none()); +}