#[cfg(test)]
mod tests;
#[cfg(test)]
mod test_attrs;
use proc_macro::TokenStream;
use quote::quote;
use syn::{ItemFn, Result, parse_macro_input};
use crate::utils::bail;
#[cfg_attr(test, mutants::skip)] pub(crate) fn enrich_err(args: TokenStream, input: TokenStream) -> TokenStream {
let args = proc_macro2::TokenStream::from(args);
let input = parse_macro_input!(input as ItemFn);
impl_enrich_err_attribute(args, input)
.unwrap_or_else(|err| err.to_compile_error())
.into()
}
fn impl_enrich_err_attribute(msg_args: proc_macro2::TokenStream, mut fn_definition: ItemFn) -> Result<proc_macro2::TokenStream> {
let msg_expr = if msg_args.is_empty() {
let msg = format!("error in function {}", &fn_definition.sig.ident);
quote! { #msg }
} else {
generate_msg_expr(msg_args)?
};
check_return_type(&fn_definition.sig.output)?;
let asyncness = &fn_definition.sig.asyncness;
let await_suffix = asyncness.is_some().then(|| quote! { .await });
let body = &fn_definition.block;
let block = quote! {
{
(#asyncness || #body)() #await_suffix .map_err(|mut e| {
let msg = #msg_expr;
ohno::Enrichable::add_enrichment(&mut e, ohno::EnrichmentEntry::new(msg, file!(), line!()));
e
})
}
};
fn_definition.block = syn::parse2(block)?;
Ok(quote! { #fn_definition })
}
pub(crate) fn generate_msg_expr(args_stream: proc_macro2::TokenStream) -> Result<proc_macro2::TokenStream> {
let tokens: Vec<_> = args_stream.into_iter().collect();
if tokens.is_empty() {
bail!("enrich_err requires a message or format string");
}
let Some(proc_macro2::TokenTree::Literal(lit)) = tokens.first() else {
bail!("cannot parse enrich_err arguments as a string literal or format expression");
};
let lit_str = lit.to_string();
if !is_quoted_string(&lit_str) {
bail!("enrich_err requires a string literal or format expression");
}
if tokens.len() > 1 || (lit_str.contains('{') && lit_str.contains('}')) {
let format_tokens = proc_macro2::TokenStream::from_iter(tokens);
Ok(quote! { format!(#format_tokens) })
} else {
Ok(quote! { #lit })
}
}
fn is_quoted_string(lit: &str) -> bool {
lit.starts_with('"') && lit.ends_with('"')
}
fn check_return_type(output: &syn::ReturnType) -> Result<()> {
match output {
syn::ReturnType::Type(_, _) => {
Ok(())
}
syn::ReturnType::Default => {
bail!("enrich_err attribute can only be applied to functions returning Result")
}
}
}
#[cfg(test)]
mod inline_tests {
use super::*;
#[test]
fn generate_msg_expr_simple() {
let expr = generate_msg_expr(quote! { "simple message" }).unwrap();
let expected = quote! { "simple message" };
assert_eq!(expr.to_string(), expected.to_string());
}
#[test]
fn generate_msg_expr_empty_args_stream() {
let err = generate_msg_expr(proc_macro2::TokenStream::new()).unwrap_err();
assert_eq!(err.to_string(), "enrich_err requires a message or format string");
}
#[test]
fn generate_msg_expr_invalid_format() {
let expr = generate_msg_expr(quote! { "simple message", 123, 345 }).unwrap();
let expected = quote! { format!("simple message", 123, 345) };
assert_eq!(expr.to_string(), expected.to_string());
}
#[test]
fn generate_msg_expr_with_one_brace() {
let expr = generate_msg_expr(quote! { "simple {message" }).unwrap();
let expected = quote! { "simple {message" };
assert_eq!(expr.to_string(), expected.to_string());
let expr = generate_msg_expr(quote! { "simple }message" }).unwrap();
let expected = quote! { "simple }message" };
assert_eq!(expr.to_string(), expected.to_string());
}
#[test]
fn generate_msg_expr_format() {
let err = generate_msg_expr(quote! { format!("error in {}", name) }).unwrap_err();
let expected_err = "cannot parse enrich_err arguments as a string literal or format expression";
assert_eq!(err.to_string(), expected_err);
}
#[test]
fn generate_msg_expr_interpolation() {
let expr = generate_msg_expr(quote! { "failed to read {path}" }).unwrap();
let expected = quote! { format!("failed to read {path}") };
assert_eq!(expr.to_string(), expected.to_string());
}
#[test]
fn test_generate_msg_expr() {
let expr = generate_msg_expr(quote! { "error in {}: {}", module, error_code }).unwrap();
let expected = quote! { format!("error in {}: {}", module, error_code) };
assert_eq!(expr.to_string(), expected.to_string());
}
#[test]
fn generate_msg_expr_multiple_tokens_no_braces() {
let expr = generate_msg_expr(quote! { "error occurred", extra_arg }).unwrap();
let expected = quote! { format!("error occurred", extra_arg) };
assert_eq!(expr.to_string(), expected.to_string());
}
#[test]
fn generate_msg_expr_invalid_literal() {
let err = generate_msg_expr(quote! { 42 }).unwrap_err();
let expected_err = "enrich_err requires a string literal or format expression";
assert_eq!(err.to_string(), expected_err);
}
#[test]
fn generate_msg_expr_boolean_literal() {
let err = generate_msg_expr(quote! { true }).unwrap_err();
let expected_err = "cannot parse enrich_err arguments as a string literal or format expression";
assert_eq!(err.to_string(), expected_err);
}
#[test]
fn generate_msg_expr_char_literal() {
let err = generate_msg_expr(quote! { 'c' }).unwrap_err();
let expected_err = "enrich_err requires a string literal or format expression";
assert_eq!(err.to_string(), expected_err);
}
#[test]
fn test_is_quoted_string() {
assert!(is_quoted_string("\"valid string\""));
assert!(is_quoted_string("\"12345\""));
assert!(is_quoted_string("\"true\""));
assert!(!is_quoted_string("\"invalid string"));
assert!(!is_quoted_string("invalid string\""));
assert!(!is_quoted_string("12345"));
assert!(!is_quoted_string("true"));
assert!(!is_quoted_string("'single quoted'"));
}
#[test]
fn check_return_type_with_result() {
let return_type: syn::ReturnType = syn::parse_quote! { -> Result<(), String> };
check_return_type(&return_type).unwrap();
}
#[test]
fn check_return_type_without_result() {
let return_type = syn::ReturnType::Default;
let err = check_return_type(&return_type).unwrap_err();
assert_eq!(
err.to_string(),
"enrich_err attribute can only be applied to functions returning Result"
);
}
}