use std::{env, fmt};
use tracing::{Dispatch, Event, Subscriber, level_filters::LevelFilter};
use tracing_subscriber::{
EnvFilter, FmtSubscriber,
field::RecordFields,
fmt::{FmtContext, FormatEvent, FormatFields, TestWriter, format, format::Writer},
registry::LookupSpan,
};
use crate::decorators::{DecorateTest, TestFn};
#[derive(Debug)]
enum Either<L, R> {
Left(L),
Right(R),
}
impl<'w, L, R> FormatFields<'w> for Either<L, R>
where
L: FormatFields<'w>,
R: FormatFields<'w>,
{
fn format_fields<F: RecordFields>(&self, writer: Writer<'w>, fields: F) -> fmt::Result {
match self {
Self::Left(formatter) => formatter.format_fields(writer, fields),
Self::Right(formatter) => formatter.format_fields(writer, fields),
}
}
}
impl<S, N, L, R> FormatEvent<S, N> for Either<L, R>
where
S: Subscriber + for<'a> LookupSpan<'a>,
N: for<'a> FormatFields<'a> + 'static,
L: FormatEvent<S, N>,
R: FormatEvent<S, N>,
{
fn format_event(
&self,
ctx: &FmtContext<'_, S, N>,
writer: Writer<'_>,
event: &Event<'_>,
) -> fmt::Result {
match self {
Self::Left(formatter) => formatter.format_event(ctx, writer, event),
Self::Right(formatter) => formatter.format_event(ctx, writer, event),
}
}
}
type TestSubscriber = FmtSubscriber<
Either<format::Pretty, format::DefaultFields>,
Either<format::Format<format::Pretty>, format::Format>,
EnvFilter,
TestWriter,
>;
#[cfg_attr(docsrs, doc(cfg(feature = "tracing")))]
#[derive(Debug, Clone, Copy)]
pub struct Trace {
directives: Option<&'static str>,
pretty: bool,
global: bool,
}
impl Trace {
pub const fn new(directives: &'static str) -> Self {
Self {
directives: Some(directives),
pretty: false,
global: false,
}
}
#[must_use]
pub const fn pretty(mut self) -> Self {
self.pretty = true;
self
}
#[must_use]
pub const fn global(mut self) -> Self {
self.global = true;
self
}
pub fn create_subscriber(self) -> impl Subscriber + for<'a> LookupSpan<'a> {
self.create_subscriber_inner()
}
fn create_subscriber_inner(self) -> TestSubscriber {
let env = env::var("RUST_LOG").ok();
let env = env.as_deref().or(self.directives).unwrap_or_default();
let env_filter = EnvFilter::builder()
.with_default_directive(LevelFilter::INFO.into())
.parse_lossy(env);
FmtSubscriber::builder()
.with_test_writer()
.with_env_filter(env_filter)
.fmt_fields(if self.pretty {
Either::Left(format::Pretty::default())
} else {
Either::Right(format::DefaultFields::default())
})
.map_event_format(|fmt| {
if self.pretty {
Either::Left(fmt.pretty())
} else {
Either::Right(fmt)
}
})
.finish()
}
}
impl<R> DecorateTest<R> for Trace {
fn decorate_and_test<F: TestFn<R>>(&'static self, test_fn: F) -> R {
let subscriber = self.create_subscriber_inner();
let _guard = if self.global {
if tracing::subscriber::set_global_default(subscriber).is_err() {
let is_test_subscriber =
tracing::dispatcher::get_default(Dispatch::is::<TestSubscriber>);
if !is_test_subscriber {
tracing::warn!("could not set up global tracing subscriber");
}
}
None
} else {
Some(tracing::subscriber::set_default(subscriber))
};
test_fn()
}
}