use crate::prelude::*;
#[cfg(not(feature = "testing"))]
use std::io::stderr;
use tracing::Subscriber;
use tracing::subscriber::set_global_default;
use tracing_subscriber::filter::Targets;
#[cfg(feature = "testing")]
use tracing_subscriber::fmt::TestWriter;
use tracing_subscriber::fmt::layer;
use tracing_subscriber::fmt::writer::BoxMakeWriter;
use tracing_subscriber::layer::SubscriberExt;
use tracing_subscriber::{Layer, Registry};
impl Logger {
pub(super) fn set(&self) -> Result<(), Report<InitError>> {
let filter = self.build_targets();
let registry = build_registry(filter, self.writer.clone());
set_global_default(registry).change_context(InitError::Init)
}
}
fn build_registry(filter: Targets, writer: Option<SharedWriter>) -> impl Subscriber {
let make_writer = writer.unwrap_or_else(|| SharedWriter::from_box(default_make_writer()));
let layer = layer()
.compact()
.with_ansi_sanitization(false)
.with_writer(make_writer)
.with_target(false)
.with_timer(ElapsedTime::default())
.with_filter(filter);
Registry::default().with(layer)
}
#[cfg(not(feature = "testing"))]
fn default_make_writer() -> BoxMakeWriter {
BoxMakeWriter::new(stderr)
}
#[cfg(feature = "testing")]
fn default_make_writer() -> BoxMakeWriter {
BoxMakeWriter::new(TestWriter::default())
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{self, Write};
use std::sync::{Arc, Mutex, MutexGuard};
use tracing::subscriber::with_default;
use tracing_subscriber::filter::LevelFilter;
use tracing_subscriber::fmt::MakeWriter;
#[derive(Clone)]
struct BufferWriter(Arc<Mutex<Vec<u8>>>);
struct BufferGuard<'a>(MutexGuard<'a, Vec<u8>>);
impl Write for BufferGuard<'_> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.0.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl<'a> MakeWriter<'a> for BufferWriter {
type Writer = BufferGuard<'a>;
fn make_writer(&'a self) -> Self::Writer {
BufferGuard(self.0.lock().expect("buffer lock"))
}
}
#[test]
fn logger_set_routes_through_custom_writer() {
let buffer = Arc::new(Mutex::new(Vec::<u8>::new()));
let writer = BufferWriter(buffer.clone());
let filter = Targets::new().with_default(LevelFilter::from(LogLevel::Info));
let registry = build_registry(filter, Some(SharedWriter::new(writer)));
with_default(registry, || {
tracing::info!("hello world");
});
let captured =
String::from_utf8(buffer.lock().expect("buffer lock").clone()).expect("utf-8");
assert!(
captured.contains("hello world"),
"captured output should contain message, got: {captured}"
);
}
}