use crate::LogRecord;
use crate::Metrics;
use crate::support::processing::RateLimiter;
use crate::validation::sanitize::LogSanitizer;
use crossbeam_channel::Sender;
use parking_lot::Mutex;
use serde_json::Value;
use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use tracing::{Event, Subscriber};
use tracing_subscriber::Layer;
use tracing_subscriber::layer::Context;
const DEFAULT_SEND_TIMEOUT_MS: u64 = 100;
const FALLBACK_BUFFER_SIZE: usize = 100;
const ERROR_SAMPLING_RATE: u64 = 100;
pub struct LoggerSubscriber {
console_sender: Sender<Arc<LogRecord>>,
async_sender: Sender<Arc<LogRecord>>,
extra_async_senders: Vec<Sender<Arc<LogRecord>>>,
metrics: Arc<Metrics>,
send_timeout_ms: u64,
fallback_buffer: Arc<Mutex<VecDeque<Arc<LogRecord>>>>,
sanitizer: Option<Arc<LogSanitizer>>,
rate_limiter: Option<Arc<RateLimiter>>,
error_sample_counter: AtomicU64,
}
impl LoggerSubscriber {
pub fn new(
console_sender: Sender<Arc<LogRecord>>,
async_sender: Sender<Arc<LogRecord>>,
metrics: Arc<Metrics>,
) -> Self {
Self {
console_sender,
async_sender,
extra_async_senders: Vec::new(),
metrics,
send_timeout_ms: DEFAULT_SEND_TIMEOUT_MS,
fallback_buffer: Arc::new(Mutex::new(VecDeque::with_capacity(FALLBACK_BUFFER_SIZE))),
sanitizer: None,
rate_limiter: None,
error_sample_counter: AtomicU64::new(0),
}
}
pub fn with_extra_async_sender(mut self, sender: Sender<Arc<LogRecord>>) -> Self {
self.extra_async_senders.push(sender);
self
}
fn send_to_async_sinks(&self, record: &Arc<LogRecord>, timeout: Duration) -> bool {
let mut all_ok = self
.async_sender
.send_timeout(Arc::clone(record), timeout)
.is_ok();
for sender in &self.extra_async_senders {
if sender.send_timeout(Arc::clone(record), timeout).is_err() {
all_ok = false;
}
}
all_ok
}
pub fn with_timeout(mut self, timeout_ms: u64) -> Self {
self.send_timeout_ms = timeout_ms;
self
}
pub fn with_sanitizer(mut self, sanitizer: Arc<LogSanitizer>) -> Self {
self.sanitizer = Some(sanitizer);
self
}
pub fn with_rate_limiter(mut self, rate_limiter: Arc<RateLimiter>) -> Self {
self.rate_limiter = Some(rate_limiter);
self
}
fn is_critical_level(level: &str) -> bool {
level == "ERROR" || level == "FATAL"
}
fn sanitize_record(&self, record: &mut LogRecord) {
if let Some(ref sanitizer) = self.sanitizer {
record.message = sanitizer.sanitize(&record.message);
for value in record.fields.values_mut() {
if let Value::String(s) = value {
*s = sanitizer.sanitize(s);
}
}
}
}
pub fn try_flush_fallback(&self) {
let mut buffer = self.fallback_buffer.lock();
while let Some(record) = buffer.front() {
let timeout = Duration::from_millis(self.send_timeout_ms);
if !self.send_to_async_sinks(record, timeout) {
break;
}
buffer.pop_front();
}
}
}
impl Drop for LoggerSubscriber {
fn drop(&mut self) {
let buffer_len = {
let buffer = self.fallback_buffer.lock();
buffer.len()
};
if buffer_len > 0 {
self.try_flush_fallback();
let remaining = self.fallback_buffer.lock().len();
if remaining > 0 {
tracing::warn!(
unflushed_records = remaining,
"LoggerSubscriber dropped with unflushed fallback records"
);
}
}
}
}
impl<S> Layer<S> for LoggerSubscriber
where
S: Subscriber,
{
fn on_event(&self, event: &Event<'_>, _ctx: Context<'_, S>) {
let mut record = LogRecord::from_event(event);
if let Some(ref limiter) = self.rate_limiter
&& !limiter.try_acquire()
{
if Self::is_critical_level(&record.level) {
let count = self.error_sample_counter.fetch_add(1, Ordering::Relaxed);
if !count.is_multiple_of(ERROR_SAMPLING_RATE) {
self.metrics.inc_logs_dropped();
return;
}
} else {
self.metrics.inc_logs_dropped();
return;
}
}
self.sanitize_record(&mut record);
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);
if !self.send_to_async_sinks(&record, timeout) {
if Self::is_critical_level(&record.level) {
let mut buffer = self.fallback_buffer.lock();
if buffer.len() >= FALLBACK_BUFFER_SIZE {
buffer.pop_front();
}
buffer.push_back(record);
} else {
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()
.push_back(Arc::clone(&record));
subscriber.try_flush_fallback();
assert!(
subscriber.fallback_buffer.lock().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()
.push_back(Arc::clone(&record));
drop(_async_rx);
subscriber.try_flush_fallback();
assert_eq!(
subscriber.fallback_buffer.lock().len(),
1,
"buffer should still contain the record after disconnect"
);
}
#[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_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,
"non-critical async timeout should NOT 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();
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"
);
}
#[test]
fn test_with_sanitizer_escapes_newline_in_message() {
let (console_tx, console_rx) = bounded(10);
let (async_tx, _async_rx) = bounded(10);
let metrics = Arc::new(Metrics::new());
let sanitizer = Arc::new(LogSanitizer::new());
let layer = LoggerSubscriber::new(console_tx, async_tx, metrics).with_sanitizer(sanitizer);
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::info!(target: "test::sanitizer", message = "line1\nline2");
});
let received = console_rx.recv().unwrap();
assert!(
received.message.contains("\\n"),
"message should contain escaped newline, got: {:?}",
received.message
);
assert!(
!received.message.contains('\n'),
"message should not contain raw newline"
);
}
#[test]
fn test_without_sanitizer_message_unchanged() {
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::no_sanitizer", message = "plain message");
});
let received = console_rx.recv().unwrap();
assert_eq!(received.message, "plain message");
}
#[test]
fn test_rate_limiter_drops_non_critical_logs() {
let (console_tx, console_rx) = bounded(100);
let (async_tx, _async_rx) = bounded(100);
let metrics = Arc::new(Metrics::new());
let limiter = Arc::new(RateLimiter::new(2));
let layer =
LoggerSubscriber::new(console_tx, async_tx, metrics.clone()).with_rate_limiter(limiter);
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
for _ in 0..10 {
tracing::info!(target: "test::rate", message = "flood");
}
});
let mut count = 0;
while console_rx.try_recv().is_ok() {
count += 1;
}
assert!(
count <= 2,
"at most 2 logs should pass rate limiter, got {}",
count
);
assert!(
metrics.logs_dropped() >= 8,
"at least 8 logs should be dropped, got {}",
metrics.logs_dropped()
);
}
#[test]
fn test_rate_limiter_samples_error_on_rejection() {
let (console_tx, console_rx) = bounded(200);
let (async_tx, _async_rx) = bounded(200);
let metrics = Arc::new(Metrics::new());
let limiter = Arc::new(RateLimiter::new(1));
let layer =
LoggerSubscriber::new(console_tx, async_tx, metrics.clone()).with_rate_limiter(limiter);
let registry = tracing_subscriber::registry().with(layer);
with_default(registry, || {
tracing::info!(target: "test::rate", message = "consume token");
for _ in 0..100 {
tracing::error!(target: "test::rate", message = "error flood");
}
});
let mut error_count = 0;
while let Ok(record) = console_rx.try_recv() {
if record.level == "ERROR" {
error_count += 1;
}
}
assert!(
(1..=5).contains(&error_count),
"expected ~1 sampled ERROR through rate limiter, got {}",
error_count
);
}
}