use crate::LogRecord;
use crate::Metrics;
use crossbeam_channel::Sender;
use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::Duration;
use tracing::{Event, Subscriber};
use tracing_subscriber::layer::Context;
use tracing_subscriber::Layer;
const DEFAULT_SEND_TIMEOUT_MS: u64 = 100;
const FALLBACK_BUFFER_SIZE: usize = 100;
pub struct LoggerSubscriber {
console_sender: Sender<Arc<LogRecord>>,
async_sender: Sender<Arc<LogRecord>>,
metrics: Arc<Metrics>,
send_timeout_ms: u64,
fallback_buffer: Arc<Mutex<VecDeque<Arc<LogRecord>>>>,
}
impl LoggerSubscriber {
pub fn new(
console_sender: Sender<Arc<LogRecord>>,
async_sender: Sender<Arc<LogRecord>>,
metrics: Arc<Metrics>,
) -> Self {
Self {
console_sender,
async_sender,
metrics,
send_timeout_ms: DEFAULT_SEND_TIMEOUT_MS,
fallback_buffer: Arc::new(Mutex::new(VecDeque::with_capacity(FALLBACK_BUFFER_SIZE))),
}
}
pub fn with_timeout(mut self, timeout_ms: u64) -> Self {
self.send_timeout_ms = timeout_ms;
self
}
fn is_critical_level(level: &str) -> bool {
level == "ERROR" || level == "FATAL"
}
pub fn try_flush_fallback(&self) {
let mut buffer = match self.fallback_buffer.lock() {
Ok(guard) => guard,
Err(poisoned) => {
tracing::warn!("Fallback buffer mutex poisoned, recovering");
poisoned.into_inner()
}
};
while let Some(record) = buffer.front() {
let timeout = Duration::from_millis(self.send_timeout_ms);
match self.async_sender.send_timeout(Arc::clone(record), timeout) {
Ok(_) => {
buffer.pop_front();
}
Err(_) => break,
}
}
}
}
impl<S> Layer<S> for LoggerSubscriber
where
S: Subscriber,
{
fn on_event(&self, event: &Event<'_>, _ctx: Context<'_, S>) {
let record = LogRecord::from_event(event);
let record = Arc::new(record);
match self.console_sender.try_send(Arc::clone(&record)) {
Ok(_) => {}
Err(crossbeam_channel::TrySendError::Full(_)) => {
self.metrics.inc_channel_blocked();
self.metrics.inc_logs_dropped();
}
Err(crossbeam_channel::TrySendError::Disconnected(_)) => {
self.metrics.inc_logs_dropped();
}
}
let timeout = Duration::from_millis(self.send_timeout_ms);
match self.async_sender.send_timeout(Arc::clone(&record), timeout) {
Ok(_) => {}
Err(crossbeam_channel::SendTimeoutError::Timeout(_)) => {
if Self::is_critical_level(&record.level) {
let mut buffer = match self.fallback_buffer.lock() {
Ok(guard) => guard,
Err(poisoned) => {
tracing::warn!("Fallback buffer mutex poisoned, recovering");
poisoned.into_inner()
}
};
if buffer.len() >= FALLBACK_BUFFER_SIZE {
buffer.pop_front();
}
buffer.push_back(record);
} else {
self.metrics.inc_channel_blocked();
self.metrics.inc_logs_dropped();
}
}
Err(crossbeam_channel::SendTimeoutError::Disconnected(_)) => {
self.metrics.inc_logs_dropped();
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crossbeam_channel::bounded;
use serial_test::serial;
use tracing::subscriber::with_default;
use tracing_subscriber::prelude::*;
#[test]
fn test_on_event_sends_to_channels() {
let (console_tx, console_rx) = bounded(10);
let (async_tx, async_rx) = bounded(10);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics);
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::info!(target: "test::subscriber", message = "hello", user_id = 1u64);
});
let console_received = console_rx.recv().unwrap();
assert_eq!(console_received.level, "INFO");
assert_eq!(console_received.target, "test::subscriber");
assert_eq!(console_received.message, "hello");
let async_received = async_rx.recv().unwrap();
assert_eq!(async_received.level, "INFO");
assert_eq!(async_received.target, "test::subscriber");
assert_eq!(async_received.message, "hello");
}
#[test]
fn test_on_event_handles_full_channel() {
let (console_tx, console_rx) = bounded(1);
let (async_tx, async_rx) = bounded(1);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics);
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
for i in 0..5 {
tracing::info!(target: "test::subscriber", message = "msg {}", i);
}
});
while console_rx.try_recv().is_ok() {}
while async_rx.try_recv().is_ok() {}
}
#[test]
fn test_critical_level_adds_to_fallback_buffer() {
let (console_tx, _console_rx) = bounded(10);
let (async_tx, _async_rx) = bounded(0);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics.clone());
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::error!(target: "test::subscriber", message = "critical error");
});
assert_eq!(metrics.logs_written(), 0);
}
#[test]
fn test_fallback_buffer_does_not_panic_on_overflow() {
let (console_tx, _cr) = bounded(10);
let (async_tx, _ar) = bounded(0);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics);
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
for i in 0..105 {
tracing::error!(target: "test::subscriber", msg = "overflow {}", i);
}
});
}
#[test]
fn test_try_flush_fallback_with_disconnected_channel() {
let (console_tx1, _cr1) = bounded(10);
let (async_tx1, _ar1) = bounded(1);
drop(_ar1);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx1.clone(), async_tx1.clone(), metrics.clone());
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::error!(target: "test::subscriber", msg = "fallback before disconnect");
});
let _subscriber_b = LoggerSubscriber::new(console_tx1, async_tx1, metrics);
_subscriber_b.try_flush_fallback();
}
#[test]
fn test_on_event_dropped_on_disconnected_async_channel() {
let (console_tx, _cr) = bounded(10);
let (async_tx, _ar) = bounded(1);
drop(_ar); let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics.clone());
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::info!(target: "test::subscriber", message = "after disconnect");
});
assert_eq!(metrics.logs_dropped(), 1);
}
#[test]
fn test_with_timeout_configures_send_timeout() {
let (console_tx, _) = bounded(10);
let (async_tx, _) = bounded(10);
let metrics = Arc::new(Metrics::new());
let subscriber = LoggerSubscriber::new(console_tx, async_tx, metrics).with_timeout(500);
assert_eq!(subscriber.send_timeout_ms, 500);
}
#[test]
fn test_try_flush_fallback_drains_buffer_on_success() {
let (console_tx, _console_rx) = bounded(10);
let (async_tx, async_rx) = bounded(10);
let metrics = Arc::new(Metrics::new());
let subscriber = LoggerSubscriber::new(console_tx, async_tx, metrics);
let record = Arc::new(LogRecord::new(
tracing::Level::ERROR,
"test::fallback".to_string(),
"fallback flush test".to_string(),
));
subscriber
.fallback_buffer
.lock()
.unwrap()
.push_back(Arc::clone(&record));
subscriber.try_flush_fallback();
assert!(
subscriber.fallback_buffer.lock().unwrap().is_empty(),
"buffer should be empty after successful flush"
);
let received = async_rx.recv_timeout(std::time::Duration::from_millis(100));
assert!(received.is_ok(), "should receive the flushed record");
assert_eq!(received.unwrap().message, "fallback flush test");
}
#[test]
fn test_try_flush_fallback_breaks_on_disconnected_channel() {
let (console_tx, _console_rx) = bounded(10);
let (async_tx, _async_rx) = bounded(10);
let metrics = Arc::new(Metrics::new());
let subscriber = LoggerSubscriber::new(console_tx, async_tx, metrics);
let record = Arc::new(LogRecord::new(
tracing::Level::ERROR,
"test::fallback".to_string(),
"disconnect test".to_string(),
));
subscriber
.fallback_buffer
.lock()
.unwrap()
.push_back(Arc::clone(&record));
drop(_async_rx);
subscriber.try_flush_fallback();
assert_eq!(
subscriber.fallback_buffer.lock().unwrap().len(),
1,
"buffer should still contain the record after disconnect"
);
}
#[test]
fn test_try_flush_fallback_recovers_from_poisoned_mutex() {
let (console_tx, _console_rx) = bounded(10);
let (async_tx, _async_rx) = bounded(10);
let metrics = Arc::new(Metrics::new());
let subscriber = LoggerSubscriber::new(console_tx, async_tx, metrics);
let buffer_clone = Arc::clone(&subscriber.fallback_buffer);
let handle = std::thread::spawn(move || {
let _guard = buffer_clone.lock().unwrap();
panic!("intentional panic to poison mutex");
});
let join_result = handle.join();
assert!(join_result.is_err(), "thread should have panicked");
subscriber.try_flush_fallback();
}
#[test]
fn test_on_event_console_disconnected_increments_dropped() {
let (console_tx, _console_rx) = bounded(10);
drop(_console_rx);
let (async_tx, _async_rx) = bounded(10);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics.clone());
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::info!(target: "test::subscriber", message = "console disconnected");
});
assert_eq!(
metrics.logs_dropped(),
1,
"console disconnect should increment logs_dropped by 1"
);
}
#[test]
fn test_on_event_console_full_channel_increments_blocked_and_dropped() {
let (console_tx, console_rx) = bounded(1);
let (async_tx, _async_rx) = bounded(10);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics.clone());
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::info!(target: "test::subscriber", message = "first");
tracing::info!(target: "test::subscriber", message = "second");
});
while console_rx.try_recv().is_ok() {}
assert!(
metrics.logs_dropped() >= 1,
"console full should increment logs_dropped, got: {}",
metrics.logs_dropped()
);
}
#[test]
fn test_on_event_critical_level_recovers_from_poisoned_fallback_mutex() {
let (console_tx, _console_rx) = bounded(10);
let (async_tx, _async_rx) = bounded(0);
let layer = LoggerSubscriber::new(console_tx, async_tx, Arc::new(Metrics::new()));
let buffer_clone = Arc::clone(&layer.fallback_buffer);
let handle = std::thread::spawn(move || {
let _guard = buffer_clone.lock().unwrap();
panic!("intentional panic to poison fallback buffer mutex");
});
let join_result = handle.join();
assert!(
join_result.is_err(),
"poisoning thread should have panicked"
);
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::error!(target: "test::subscriber", message = "poisoned fallback test");
});
let console_received = _console_rx.try_recv();
assert!(
console_received.is_ok(),
"console channel should still receive the record"
);
assert_eq!(console_received.unwrap().level, "ERROR");
}
#[test]
fn test_on_event_console_ok_and_async_ok_paths() {
let (console_tx, console_rx) = bounded(10);
let (async_tx, async_rx) = bounded(10);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics.clone());
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::info!(target: "test::subscriber", message = "ok path test");
});
assert!(
console_rx.try_recv().is_ok(),
"console should receive record"
);
assert!(async_rx.try_recv().is_ok(), "async should receive record");
assert_eq!(
metrics.logs_dropped(),
0,
"Ok path should not increment logs_dropped"
);
}
#[test]
#[serial]
fn test_on_event_non_critical_async_timeout_increments_blocked_and_dropped() {
let (console_tx, _console_rx) = bounded(10);
let (async_tx, _async_rx) = bounded(0);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics.clone());
let registry = tracing_subscriber::registry().with(layer);
let before_blocked = metrics.channel_blocked();
let before_dropped = metrics.logs_dropped();
with_default(registry, || {
tracing::info!(target: "test::subscriber", message = "non-critical timeout");
});
assert_eq!(
metrics.channel_blocked(),
before_blocked + 1,
"non-critical async timeout should increment channel_blocked"
);
assert_eq!(
metrics.logs_dropped(),
before_dropped + 1,
"non-critical async timeout should increment logs_dropped"
);
}
#[test]
#[serial]
fn test_on_event_critical_async_timeout_stores_record_in_fallback_buffer() {
let (console_tx, _console_rx) = bounded(10);
let (async_tx, _async_rx) = bounded(0);
let metrics = Arc::new(Metrics::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics.clone());
let fallback_buffer = Arc::clone(&layer.fallback_buffer);
let registry = tracing_subscriber::registry().with(layer);
let before_blocked = metrics.channel_blocked();
let before_dropped = metrics.logs_dropped();
with_default(registry, || {
tracing::error!(target: "test::subscriber", message = "critical timeout");
});
let buffer_guard = fallback_buffer.lock().unwrap();
assert_eq!(
buffer_guard.len(),
1,
"fallback_buffer should contain exactly 1 record"
);
let record = buffer_guard
.front()
.expect("should have a record in fallback_buffer");
assert_eq!(record.level, "ERROR", "record level should be ERROR");
assert_eq!(
record.message, "critical timeout",
"record message should match"
);
drop(buffer_guard);
assert_eq!(
metrics.channel_blocked(),
before_blocked,
"critical level should not increment channel_blocked"
);
assert_eq!(
metrics.logs_dropped(),
before_dropped,
"critical level should not increment logs_dropped"
);
}
}