use anyhow::{Context, Result};
use async_nats::Client as NatsClient;
use chrono::{DateTime, Utc};
use governor::{
clock::DefaultClock, middleware::NoOpMiddleware, state::direct::NotKeyed, state::InMemoryState,
Quota, RateLimiter,
};
use serde::{Deserialize, Serialize};
use smith_bus::subjects::builders::LogSubject;
use smith_config::{LoggingConfig, NatsConfig, NatsLoggingConfig};
use std::collections::HashMap;
use std::num::NonZeroU32;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::mpsc;
use tokio::time::interval;
use tracing::{Event, Subscriber};
use tracing_subscriber::{
layer::{Context as TracingContext, SubscriberExt},
util::SubscriberInitExt,
EnvFilter, Layer, Registry,
};
use uuid::Uuid;
pub mod error;
pub mod metrics;
pub use error::{LoggingError, LoggingResult};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LogEntry {
pub timestamp: DateTime<Utc>,
pub level: String,
pub service: String,
pub target: String,
pub message: String,
pub fields: HashMap<String, serde_json::Value>,
pub span: Option<SpanInfo>,
pub trace: Option<TraceInfo>,
pub correlation_id: String,
pub node_id: String,
pub metadata: LogMetadata,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SpanInfo {
pub id: String,
pub parent_id: Option<String>,
pub name: String,
pub fields: HashMap<String, serde_json::Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TraceInfo {
pub id: String,
pub context: HashMap<String, String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LogMetadata {
pub file: Option<String>,
pub line: Option<u32>,
pub module_path: Option<String>,
pub thread_id: Option<String>,
pub performance_category: Option<String>,
}
pub struct NatsLoggingLayer {
service_name: String,
config: NatsLoggingConfig,
log_sender: mpsc::UnboundedSender<LogEntry>,
rate_limiter: Option<Arc<RateLimiter<NotKeyed, InMemoryState, DefaultClock, NoOpMiddleware>>>,
node_id: String,
}
struct LogProcessor {
nats_client: NatsClient,
config: NatsLoggingConfig,
log_receiver: mpsc::UnboundedReceiver<LogEntry>,
buffer: Vec<LogEntry>,
service_name: String,
}
pub struct LoggingGuard {
_processor_handle: tokio::task::JoinHandle<Result<()>>,
}
impl Drop for LoggingGuard {
fn drop(&mut self) {
tracing::debug!("Logging infrastructure shutting down");
}
}
impl NatsLoggingLayer {
pub fn new(
service_name: String,
config: NatsLoggingConfig,
nats_client: NatsClient,
) -> Result<(Self, LoggingGuard)> {
let (log_sender, log_receiver) = mpsc::unbounded_channel();
let rate_limiter = if config.rate_limit > 0 {
let quota = Quota::per_second(
NonZeroU32::new(config.rate_limit as u32)
.context("Invalid rate limit configuration")?,
);
Some(Arc::new(RateLimiter::direct(quota)))
} else {
None
};
let short_uuid = Uuid::new_v4().to_string();
let node_id = format!(
"{}_{}",
hostname::get().unwrap_or_default().to_string_lossy(),
&short_uuid[..8]
);
let processor = LogProcessor {
nats_client,
config: config.clone(),
log_receiver,
buffer: Vec::with_capacity(config.batch_size),
service_name: service_name.clone(),
};
let processor_handle = tokio::spawn(async move { processor.run().await });
let layer = Self {
service_name,
config,
log_sender,
rate_limiter,
node_id,
};
let guard = LoggingGuard {
_processor_handle: processor_handle,
};
Ok((layer, guard))
}
fn should_process(&self, event: &Event) -> bool {
if let Some(ref level_filter) = self.config.level_filter {
let event_level = event.metadata().level();
let filter_level = match level_filter.as_str() {
"error" => tracing::Level::ERROR,
"warn" => tracing::Level::WARN,
"info" => tracing::Level::INFO,
"debug" => tracing::Level::DEBUG,
"trace" => tracing::Level::TRACE,
_ => return true, };
if *event_level > filter_level {
return false;
}
}
if !self.config.target_filters.is_empty() {
let target = event.metadata().target();
let matches = self
.config
.target_filters
.iter()
.any(|filter| target.starts_with(filter));
if !matches {
return false;
}
}
if let Some(ref rate_limiter) = self.rate_limiter {
if rate_limiter.check().is_err() {
return false;
}
}
true
}
fn event_to_log_entry<S>(&self, event: &Event, ctx: &TracingContext<'_, S>) -> LogEntry
where
S: Subscriber + for<'a> tracing_subscriber::registry::LookupSpan<'a>,
{
let metadata = event.metadata();
let mut field_visitor = FieldVisitor::new();
event.record(&mut field_visitor);
let span = if self.config.include_spans {
ctx.event_span(event).map(|span_ref| {
let span_metadata = span_ref.metadata();
let mut span_fields = HashMap::new();
let span_name = span_metadata.name().to_string();
span_fields.insert(
"span_name".to_string(),
serde_json::Value::String(span_name),
);
SpanInfo {
id: format!("{:x}", span_ref.id().into_u64()),
parent_id: span_ref
.parent()
.map(|p| format!("{:x}", p.id().into_u64())),
name: span_metadata.name().to_string(),
fields: span_fields,
}
})
} else {
None
};
let correlation_id = Uuid::new_v4().to_string();
let log_metadata = LogMetadata {
file: if cfg!(debug_assertions) {
metadata.file().map(|s| s.to_string())
} else {
None
},
line: if cfg!(debug_assertions) {
metadata.line()
} else {
None
},
module_path: metadata.module_path().map(|s| s.to_string()),
thread_id: Some(format!("{:?}", std::thread::current().id())),
performance_category: field_visitor
.fields
.get("performance_category")
.and_then(|v| v.as_str().map(|s| s.to_string())),
};
LogEntry {
timestamp: Utc::now(),
level: metadata.level().to_string().to_uppercase(),
service: self.service_name.clone(),
target: metadata.target().to_string(),
message: field_visitor.message,
fields: field_visitor.fields,
span,
trace: None, correlation_id,
node_id: self.node_id.clone(),
metadata: log_metadata,
}
}
}
impl<S> Layer<S> for NatsLoggingLayer
where
S: Subscriber + for<'a> tracing_subscriber::registry::LookupSpan<'a>,
{
fn on_event(&self, event: &Event<'_>, ctx: TracingContext<'_, S>) {
if !self.should_process(event) {
return;
}
let log_entry = self.event_to_log_entry(event, &ctx);
if self.log_sender.send(log_entry).is_err() {
if self.config.fallback_to_console {
eprintln!(
"NATS logging unavailable, log entry lost: {}",
event.metadata().target()
);
}
}
}
}
struct FieldVisitor {
message: String,
fields: HashMap<String, serde_json::Value>,
}
impl FieldVisitor {
fn new() -> Self {
Self {
message: String::new(),
fields: HashMap::new(),
}
}
}
impl tracing::field::Visit for FieldVisitor {
fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn std::fmt::Debug) {
let value_str = format!("{:?}", value);
if field.name() == "message" {
self.message = value_str;
} else {
self.fields.insert(
field.name().to_string(),
serde_json::Value::String(value_str),
);
}
}
fn record_str(&mut self, field: &tracing::field::Field, value: &str) {
if field.name() == "message" {
self.message = value.to_string();
} else {
self.fields.insert(
field.name().to_string(),
serde_json::Value::String(value.to_string()),
);
}
}
fn record_i64(&mut self, field: &tracing::field::Field, value: i64) {
self.fields.insert(
field.name().to_string(),
serde_json::Value::Number(value.into()),
);
}
fn record_u64(&mut self, field: &tracing::field::Field, value: u64) {
self.fields.insert(
field.name().to_string(),
serde_json::Value::Number(value.into()),
);
}
fn record_bool(&mut self, field: &tracing::field::Field, value: bool) {
self.fields
.insert(field.name().to_string(), serde_json::Value::Bool(value));
}
}
impl LogProcessor {
async fn run(mut self) -> Result<()> {
let mut batch_timer = if self.config.batch_enabled {
Some(interval(Duration::from_millis(
self.config.batch_timeout_ms,
)))
} else {
None
};
loop {
tokio::select! {
log_entry = self.log_receiver.recv() => {
match log_entry {
Some(entry) => {
if self.config.batch_enabled {
self.buffer.push(entry);
if self.buffer.len() >= self.config.batch_size {
self.flush_buffer().await?;
}
} else {
self.publish_single(entry).await?;
}
}
None => {
if !self.buffer.is_empty() {
self.flush_buffer().await?;
}
break;
}
}
}
_ = async {
if let Some(ref mut timer) = batch_timer {
timer.tick().await;
} else {
std::future::pending::<()>().await;
}
} => {
if !self.buffer.is_empty() {
self.flush_buffer().await?;
}
}
}
}
Ok(())
}
async fn publish_single(&self, entry: LogEntry) -> Result<()> {
let subject = self.get_subject_for_entry(&entry);
let payload = serde_json::to_vec(&entry).context("Failed to serialize log entry")?;
let timeout = Duration::from_millis(self.config.publish_timeout_ms);
for attempt in 1..=self.config.max_retries {
match tokio::time::timeout(
timeout,
self.nats_client
.publish(subject.clone(), payload.clone().into()),
)
.await
{
Ok(Ok(_)) => return Ok(()),
Ok(Err(e)) => {
tracing::warn!(
"NATS publish failed (attempt {}/{}): {}",
attempt,
self.config.max_retries,
e
);
if attempt < self.config.max_retries {
tokio::time::sleep(Duration::from_millis(100 * attempt as u64)).await;
}
}
Err(_) => {
tracing::warn!(
"NATS publish timeout (attempt {}/{})",
attempt,
self.config.max_retries
);
if attempt < self.config.max_retries {
tokio::time::sleep(Duration::from_millis(100 * attempt as u64)).await;
}
}
}
}
if self.config.fallback_to_console {
eprintln!(
"NATS logging failed, falling back to console: {}",
entry.message
);
}
Err(anyhow::anyhow!(
"Failed to publish log entry after {} retries",
self.config.max_retries
))
}
async fn flush_buffer(&mut self) -> Result<()> {
if self.buffer.is_empty() {
return Ok(());
}
let entries =
std::mem::replace(&mut self.buffer, Vec::with_capacity(self.config.batch_size));
let chunk_size = 10; for chunk in entries.chunks(chunk_size) {
let tasks: Vec<_> = chunk
.iter()
.map(|entry| self.publish_single(entry.clone()))
.collect();
let results = futures::future::join_all(tasks).await;
for (i, result) in results.iter().enumerate() {
if let Err(e) = result {
tracing::debug!("Failed to publish log entry {}: {}", i, e);
}
}
}
Ok(())
}
fn get_subject_for_entry(&self, entry: &LogEntry) -> String {
match entry.level.as_str() {
"ERROR" => LogSubject::error(&self.service_name),
"WARN" | "INFO" | "DEBUG" | "TRACE" => {
LogSubject::service(&self.service_name, &entry.level.to_lowercase())
}
_ => LogSubject::service(&self.service_name, &"unknown".to_string()),
}
}
}
pub async fn init_logging(
logging_config: &LoggingConfig,
nats_config: &NatsConfig,
service_name: &str,
) -> Result<Option<LoggingGuard>> {
let env_filter =
EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new(&logging_config.level));
let fmt_layer = tracing_subscriber::fmt::layer()
.with_target(true)
.with_timer(tracing_subscriber::fmt::time::ChronoUtc::rfc_3339())
.with_level(true);
let registry = Registry::default().with(env_filter).with(fmt_layer);
if logging_config.nats.enabled {
let nats_client = async_nats::connect(&nats_config.url)
.await
.context("Failed to connect to NATS for logging")?;
let (nats_layer, guard) = NatsLoggingLayer::new(
service_name.to_string(),
logging_config.nats.clone(),
nats_client,
)?;
match registry.with(nats_layer).try_init() {
Ok(()) => {
tracing::info!(
"Logging initialized with NATS integration for service: {}",
service_name
);
Ok(Some(guard))
}
Err(err) => {
tracing::warn!(
"Logging already initialized, skipping duplicate subscriber: {}",
err
);
drop(guard);
Ok(None)
}
}
} else {
match registry.try_init() {
Ok(()) => {
tracing::info!(
"Logging initialized without NATS integration for service: {}",
service_name
);
}
Err(err) => {
tracing::warn!(
"Logging already initialized, skipping duplicate subscriber: {}",
err
);
}
}
Ok(None)
}
}
pub fn init_console_logging(level: &str) -> Result<()> {
let env_filter = EnvFilter::new(level);
let builder = tracing_subscriber::fmt()
.with_env_filter(env_filter)
.with_target(true)
.with_timer(tracing_subscriber::fmt::time::ChronoUtc::rfc_3339())
.with_level(true);
let _ = builder.try_init();
Ok(())
}
pub use tracing::{debug, error, info, trace, warn};