use crate::Result;
use std::{
cell::RefCell,
future::Future,
io::Write,
panic::AssertUnwindSafe,
sync::{Arc, Mutex},
time::Duration,
};
use tracing_subscriber::fmt::MakeWriter;
#[derive(Clone)]
struct TestWriter {
log_events: Arc<Mutex<Vec<u8>>>,
}
impl TestWriter {
fn new() -> Self {
Self {
log_events: Arc::new(Mutex::new(Vec::<u8>::new())),
}
}
fn take_string(&self) -> String {
let mut guard = self.log_events.lock().unwrap();
let buffer: Vec<u8> = std::mem::take(&mut guard);
String::from_utf8(buffer).unwrap()
}
}
impl<'a> Write for &'a TestWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
let mut guard = self.log_events.lock().unwrap();
guard.write(buf)
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
impl<'a> MakeWriter<'a> for TestWriter {
type Writer = &'a Self;
fn make_writer(&'a self) -> Self::Writer {
self
}
}
pub fn test_with_logging(test: impl Future<Output = Result<()>>) -> Result<()> {
let test_writer = TestWriter::new();
let dispatch = {
use tracing_subscriber::prelude::*;
use tracing_subscriber::{fmt, EnvFilter};
let format = fmt::layer()
.with_level(true) .with_target(true) .with_thread_ids(true) .with_thread_names(false) .with_writer(test_writer.clone());
let filter = EnvFilter::try_from_default_env()
.or_else(|_| EnvFilter::try_new("h2=warn,hyper=info,rustls=info,aws=info,debug"))
.unwrap();
let subscriber = tracing_subscriber::registry().with(filter).with(format);
tracing::Dispatch::new(subscriber)
};
let dispatch = Arc::new(dispatch);
tracing::dispatcher::with_default(&dispatch, || {
std::thread_local! {
static THREAD_DISPATCHER_GUARD: RefCell<Option<tracing::subscriber::DefaultGuard>> = RefCell::new(None);
}
let mut builder = tokio::runtime::Builder::new_multi_thread();
builder.enable_all();
{
let dispatch = dispatch.clone();
builder.on_thread_start(move || {
let dispatch = dispatch.clone();
THREAD_DISPATCHER_GUARD.with(|cell| {
cell.replace(Some(tracing::dispatcher::set_default(&dispatch)));
})
});
}
builder.on_thread_stop(|| {
THREAD_DISPATCHER_GUARD.with(|cell| cell.replace(None));
});
let runtime = builder.build()?;
let result = std::panic::catch_unwind(AssertUnwindSafe(move || {
let result = runtime.block_on(test);
runtime.shutdown_timeout(Duration::from_secs(10));
result
}));
let log_events = test_writer.take_string();
println!("Log events from this test: \n{}", log_events);
match result {
Ok(result) => {
result
}
Err(err) => {
std::panic::resume_unwind(err)
}
}
})?;
Ok(())
}