1use serde_json::{Map, Value};
8use std::sync::{Arc, Mutex};
9use tokio::sync::mpsc;
10use tracing::subscriber::Interest;
11use tracing::{Level, Subscriber};
12use tracing_subscriber::Layer;
13use tracing_subscriber::filter::LevelFilter;
14use tracing_subscriber::layer::Context;
15
16#[derive(Clone, Debug)]
18pub struct LogEvent {
19 pub level: Level,
20 pub logger: String,
21 pub data: Value,
22}
23
24pub struct McpLoggingLayer {
27 event_tx: mpsc::UnboundedSender<LogEvent>,
28 log_level_filter: Arc<Mutex<LevelFilter>>,
29}
30
31impl McpLoggingLayer {
32 pub fn new(
33 event_tx: mpsc::UnboundedSender<LogEvent>,
34 log_level_filter: Arc<Mutex<LevelFilter>>,
35 ) -> Self {
36 Self {
37 event_tx,
38 log_level_filter,
39 }
40 }
41}
42
43impl<S> Layer<S> for McpLoggingLayer
44where
45 S: Subscriber,
46{
47 fn on_event(&self, event: &tracing::Event<'_>, _ctx: Context<'_, S>) {
48 let metadata = event.metadata();
49 let level = *metadata.level();
50 let target = metadata.target();
51
52 let filter_level = self
54 .log_level_filter
55 .lock()
56 .unwrap_or_else(std::sync::PoisonError::into_inner);
57 if level > *filter_level {
58 return;
59 }
60 drop(filter_level);
61
62 let mut fields = Map::new();
64 let mut visitor = MessageVisitor(&mut fields);
65 event.record(&mut visitor);
66
67 let logger = target.to_string();
68 let data = Value::Object(fields);
69
70 let log_event = LogEvent {
72 level,
73 logger,
74 data,
75 };
76
77 let _ = self.event_tx.send(log_event);
79 }
80
81 fn register_callsite(&self, metadata: &'static tracing::Metadata<'static>) -> Interest {
82 let filter_level = self
83 .log_level_filter
84 .lock()
85 .unwrap_or_else(std::sync::PoisonError::into_inner);
86 if *metadata.level() <= *filter_level {
87 Interest::always()
88 } else {
89 Interest::never()
90 }
91 }
92
93 fn enabled(&self, metadata: &tracing::Metadata<'_>, _ctx: Context<'_, S>) -> bool {
94 let filter_level = self
95 .log_level_filter
96 .lock()
97 .unwrap_or_else(std::sync::PoisonError::into_inner);
98 *metadata.level() <= *filter_level
99 }
100}
101
102struct MessageVisitor<'a>(&'a mut Map<String, Value>);
104
105impl tracing::field::Visit for MessageVisitor<'_> {
106 fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn std::fmt::Debug) {
107 self.0.insert(
108 field.name().to_string(),
109 Value::String(format!("{value:?}")),
110 );
111 }
112
113 fn record_str(&mut self, field: &tracing::field::Field, value: &str) {
114 self.0
115 .insert(field.name().to_string(), Value::String(value.to_string()));
116 }
117
118 fn record_u64(&mut self, field: &tracing::field::Field, value: u64) {
119 self.0
120 .insert(field.name().to_string(), Value::Number(value.into()));
121 }
122
123 fn record_i64(&mut self, field: &tracing::field::Field, value: i64) {
124 self.0
125 .insert(field.name().to_string(), Value::Number(value.into()));
126 }
127
128 fn record_bool(&mut self, field: &tracing::field::Field, value: bool) {
129 self.0.insert(field.name().to_string(), Value::Bool(value));
130 }
131}
132
133#[cfg(test)]
134mod tests {
135 use super::*;
136 use tracing::Level;
137 #[test]
138 fn test_log_event_level_is_tracing_level() {
139 let event = LogEvent {
140 level: Level::ERROR,
141 logger: "test".to_string(),
142 data: Value::Null,
143 };
144 assert_eq!(event.level, Level::ERROR);
145 }
146}