extern crate proc_macro;
use proc_macro::TokenStream;
use proc_macro2::{Ident, Span, TokenStream as Tokens};
use proc_macro_crate::{crate_name, FoundCrate};
use quote::quote;
use syn::parse::Parser;
use syn::punctuated::Punctuated;
use syn::{parse_macro_input, parse_quote, ItemFn, Meta, ReturnType, Token};
#[proc_macro_attribute]
pub fn test(attr: TokenStream, item: TokenStream) -> TokenStream {
let parser = Punctuated::<Meta, Token![,]>::parse_terminated;
let args = match parser.parse2(attr.into()) {
Ok(args) => args,
Err(e) => panic!("{}", e),
};
let input = parse_macro_input!(item as ItemFn);
let args = args.into_iter().collect::<Vec<_>>();
let inner_test = match args.as_slice() {
[] => parse_quote! { ::core::prelude::v1::test },
[Meta::Path(path)] => quote! {#path},
[Meta::List(list)] => {
let path = &list.path;
let args = &list.tokens;
quote! { #path(#args) }
}
_ => {
panic!("unsupported attributes supplied: {}", quote! { args })
}
};
expand_wrapper(inner_test, &input)
}
fn expand_logging_init() -> Tokens {
let found_crate = crate_name("wick-logger").expect("wick-logger needs to be added in `Cargo.toml`");
match found_crate {
FoundCrate::Itself => quote! {
let logging_options = crate::LoggingOptionsBuilder::default()
.app_name("test")
.otlp_endpoint(std::env::var("OTLP_ENDPOINT").ok())
.levels(crate::LogFilters::with_level(crate::LogLevel::Trace))
.build()
.unwrap();
let __guard = crate::init_test(&logging_options);
},
FoundCrate::Name(name) => {
let ident = Ident::new(&name, Span::call_site());
quote! {
let logging_options = #ident::LoggingOptionsBuilder::default()
.app_name("test")
.otlp_endpoint(std::env::var("OTLP_ENDPOINT").ok())
.levels(#ident::LogFilters::with_level(#ident::LogLevel::Trace))
.build()
.unwrap();
let __guard = #ident::init_test(&logging_options);
}
}
}
}
fn expand_wrapper(inner_test: Tokens, wrappee: &ItemFn) -> TokenStream {
let attrs = &wrappee.attrs;
let async_ = &wrappee.sig.asyncness;
let await_ = if async_.is_some() {
quote! {.instrument(span).await}
} else {
quote! {}
};
let enter_ = if async_.is_some() {
quote! {use tracing::Instrument;}
} else {
quote! {let _guard = span.enter();}
};
let exit_ = if async_.is_some() {
quote! {
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
}
} else {
quote! {
drop(_guard);
}
};
let body = &wrappee.block;
let test_name = &wrappee.sig.ident;
let ret = match &wrappee.sig.output {
ReturnType::Default => quote! {},
ReturnType::Type(_, type_) => quote! {-> #type_},
};
let logging_init = expand_logging_init();
let result = quote! {
#[#inner_test]
#(#attrs)*
#async_ fn #test_name() #ret {
#async_ fn test_impl() #ret {
#body
}
#logging_init
let span = tracing::info_span!(stringify!(#test_name));
#enter_
let result = test_impl()#await_;
if let Err(e) = &result {
tracing::error!(error = ?e, "test failed");
}
#exit_
if let Some(guard) = __guard { guard.teardown() } ;
result
}
};
result.into()
}