use crate::shared::MessageRegistry;
#[cfg(not(feature = "legacy-spec"))]
use crate::types::Message;
#[cfg(not(feature = "legacy-spec"))]
use crate::types::notification::LoggingLevel;
use crate::types::notification::formatter::build_notification;
use once_cell::sync::Lazy;
use std::io::{self, Write};
use tracing::{
field::Field,
span::Attributes,
{Event, Id, Subscriber, field::Visit},
};
use tracing_subscriber::{
registry::LookupSpan,
{Layer, layer::Context},
};
const MCP_SESSION_ID: &str = "mcp_session_id";
#[cfg(not(feature = "legacy-spec"))]
pub(super) const MCP_LOG_LEVEL: &str = "mcp_log_level";
pub(crate) static LOG_REGISTRY: Lazy<MessageRegistry> = Lazy::new(MessageRegistry::new);
pub fn layer() -> MpscLayer {
MpscLayer
}
#[derive(Debug, Default)]
pub struct MpscLayer;
thread_local! {
static DELIVERING: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
}
struct DeliveryGuard;
impl DeliveryGuard {
fn enter() -> Option<Self> {
(!DELIVERING.replace(true)).then_some(Self)
}
}
impl Drop for DeliveryGuard {
fn drop(&mut self) {
DELIVERING.set(false);
}
}
impl<S> Layer<S> for MpscLayer
where
S: Subscriber + for<'a> LookupSpan<'a>,
{
#[inline]
fn on_new_span(&self, attrs: &Attributes<'_>, id: &Id, ctx: Context<'_, S>) {
record_span_context(attrs, id, &ctx);
}
#[inline]
fn on_event(&self, event: &Event<'_>, ctx: Context<'_, S>) {
let Some(_delivery) = DeliveryGuard::enter() else {
return;
};
let notification = build_notification(event);
if let Some(span) = ctx.event_span(event) {
let mut notification = notification;
notification.session_id = span
.scope()
.find_map(|s| s.extensions().get::<uuid::Uuid>().cloned());
#[cfg(not(feature = "legacy-spec"))]
if notification.method.as_str() == crate::types::notification::commands::MESSAGE {
let requested = span.scope().find_map(|s| {
s.extensions()
.get::<super::formatter::MinLogSeverity>()
.map(|m| m.0)
});
let event_severity = super::formatter::notification_severity(¬ification)
.unwrap_or_else(|| LoggingLevel::from(event.metadata().level()).severity());
if !super::formatter::message_delivered(requested, event_severity) {
return;
}
}
#[cfg(not(feature = "legacy-spec"))]
if let Some(session_id) = notification.session_id
&& let Some(sink) = super::sink::REQUEST_NOTIFICATIONS.get(&session_id)
{
let _ = sink.try_send(Message::Notification(notification));
return;
}
let _ = LOG_REGISTRY.send(notification.into());
} else {
let mut stderr = io::stderr();
let json = serde_json::to_string(¬ification).unwrap();
let _ = writeln!(stderr, "{json}");
}
}
}
#[derive(Debug, Default)]
pub struct SpanContextLayer;
impl<S> Layer<S> for SpanContextLayer
where
S: Subscriber + for<'a> LookupSpan<'a>,
{
#[inline]
fn on_new_span(&self, attrs: &Attributes<'_>, id: &Id, ctx: Context<'_, S>) {
record_span_context(attrs, id, &ctx);
}
}
pub fn span_context() -> SpanContextLayer {
SpanContextLayer
}
#[inline]
fn record_span_context<S>(attrs: &Attributes<'_>, id: &Id, ctx: &Context<'_, S>)
where
S: Subscriber + for<'a> LookupSpan<'a>,
{
let mut visitor = SpanVisitor::default();
attrs.record(&mut visitor);
if let Some(span) = ctx.span(id) {
if let Some(mcp_session_id) = visitor.session_id {
span.extensions_mut().insert(mcp_session_id);
}
#[cfg(not(feature = "legacy-spec"))]
if let Some(min) = visitor.min_severity {
span.extensions_mut()
.insert(super::formatter::MinLogSeverity(min));
}
}
}
#[derive(Default)]
struct SpanVisitor {
session_id: Option<uuid::Uuid>,
#[cfg(not(feature = "legacy-spec"))]
min_severity: Option<u8>,
}
impl Visit for SpanVisitor {
#[inline]
fn record_str(&mut self, field: &Field, value: &str) {
if field.name() == MCP_SESSION_ID
&& let Ok(session_id) = uuid::Uuid::parse_str(value)
{
self.session_id = Some(session_id);
}
}
#[cfg(not(feature = "legacy-spec"))]
#[inline]
fn record_u64(&mut self, field: &Field, value: u64) {
if field.name() == MCP_LOG_LEVEL {
self.min_severity = Some(value as u8);
}
}
fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
if field.name() == MCP_SESSION_ID && self.session_id.is_none() {
let formatted = format!("{value:?}");
let stripped = formatted
.strip_prefix('"')
.and_then(|s| s.strip_suffix('"'))
.unwrap_or(&formatted);
if let Ok(session_id) = uuid::Uuid::parse_str(stripped) {
self.session_id = Some(session_id);
}
}
}
}
#[cfg(all(test, feature = "legacy-spec"))]
mod legacy_tests {
use tracing_subscriber::prelude::*;
#[test]
fn a_delivery_refuses_to_nest() {
let outer = super::DeliveryGuard::enter().expect("the first delivery proceeds");
assert!(
super::DeliveryGuard::enter().is_none(),
"what a delivery logs must not be delivered"
);
drop(outer);
assert!(
super::DeliveryGuard::enter().is_some(),
"and the next event is delivered as usual"
);
}
fn emit_and_drain(session_id: uuid::Uuid, emit: impl FnOnce()) -> Vec<serde_json::Value> {
let (tx, mut rx) = tokio::sync::mpsc::channel(8);
super::LOG_REGISTRY.register(session_id, 1, tx);
let subscriber = tracing_subscriber::registry().with(super::MpscLayer);
tracing::subscriber::with_default(subscriber, || {
let span = tracing::info_span!("request", mcp_session_id = session_id.to_string());
let _entered = span.enter();
emit();
});
super::LOG_REGISTRY.unregister(&session_id);
let mut got = Vec::new();
while let Ok(msg) = rx.try_recv() {
got.push(serde_json::to_value(&msg).unwrap());
}
got
}
#[test]
fn a_full_queue_drops_its_overflow() {
let session_id = uuid::Uuid::new_v4();
let (tx, mut rx) = tokio::sync::mpsc::channel(1);
super::LOG_REGISTRY.register(session_id, 1, tx);
let subscriber = tracing_subscriber::registry().with(super::MpscLayer);
tracing::subscriber::with_default(subscriber, || {
let span = tracing::info_span!("request", mcp_session_id = session_id.to_string());
let _entered = span.enter();
for i in 0..64 {
tracing::info!(logger = "neva", "message {i}");
}
});
super::LOG_REGISTRY.unregister(&session_id);
let mut delivered = 0;
while rx.try_recv().is_ok() {
delivered += 1;
}
assert_eq!(
delivered, 1,
"the queue holds one, and the overflow is dropped rather than queued"
);
}
#[test]
fn a_progress_report_is_on_the_stream_before_the_emitter_returns() {
let got = emit_and_drain(uuid::Uuid::new_v4(), || {
for value in [0, 50, 100] {
tracing::info!(target: "progress", token = "tok-1", value = value, total = 100);
}
});
let progress = got
.iter()
.filter(|m| m["method"] == "notifications/progress")
.collect::<Vec<_>>();
assert_eq!(progress.len(), 3, "got: {got:?}");
assert_eq!(progress[0]["params"]["progressToken"], "tok-1");
assert_eq!(progress[0]["params"]["progress"], 0.0);
assert_eq!(progress[2]["params"]["progress"], 100.0);
assert_eq!(progress[2]["params"]["total"], 100.0);
}
#[test]
fn a_log_message_travels_the_same_way() {
let got = emit_and_drain(uuid::Uuid::new_v4(), || {
tracing::warn!(logger = "tool", "something happened");
});
assert_eq!(got.len(), 1, "got: {got:?}");
assert_eq!(got[0]["method"], "notifications/message");
assert_eq!(got[0]["params"]["level"], "warning");
assert_eq!(got[0]["params"]["data"]["message"], "something happened");
}
#[test]
fn an_event_for_an_unknown_session_is_dropped_rather_than_kept() {
let subscriber = tracing_subscriber::registry().with(super::MpscLayer);
tracing::subscriber::with_default(subscriber, || {
let span =
tracing::info_span!("request", mcp_session_id = uuid::Uuid::new_v4().to_string());
let _entered = span.enter();
tracing::warn!(logger = "tool", "nobody is listening");
});
}
}
#[cfg(all(test, not(feature = "legacy-spec")))]
mod tests {
use crate::types::notification::{LoggingLevel, NotificationFormatter};
use std::io::Write;
use std::sync::{Arc, Mutex};
use tracing_subscriber::fmt::MakeWriter;
use tracing_subscriber::prelude::*;
#[derive(Clone)]
struct BufWriter(Arc<Mutex<Vec<u8>>>);
struct BufGuard(Arc<Mutex<Vec<u8>>>);
impl Write for BufGuard {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
impl<'a> MakeWriter<'a> for BufWriter {
type Writer = BufGuard;
fn make_writer(&'a self) -> Self::Writer {
BufGuard(self.0.clone())
}
}
fn emit_within_request(log_level: Option<LoggingLevel>) -> Vec<String> {
let buf = Arc::new(Mutex::new(Vec::new()));
let subscriber = tracing_subscriber::registry()
.with(super::span_context())
.with(
tracing_subscriber::fmt::layer()
.event_format(NotificationFormatter)
.with_writer(BufWriter(buf.clone())),
);
tracing::subscriber::with_default(subscriber, || {
let span = match log_level {
Some(level) => {
tracing::info_span!("request", mcp_log_level = u64::from(level.severity()))
}
None => tracing::info_span!("request"),
};
let _entered = span.enter();
tracing::error!(logger = "tool", "error message");
tracing::warn!(logger = "tool", "warning message");
tracing::info!(logger = "tool", "info message");
tracing::debug!(logger = "tool", "debug message");
});
let raw = buf.lock().unwrap().clone();
String::from_utf8(raw)
.unwrap()
.lines()
.filter(|l| !l.trim().is_empty())
.map(str::to_owned)
.collect()
}
fn levels(lines: &[String]) -> Vec<String> {
lines
.iter()
.filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
.filter(|v| v["method"] == "notifications/message")
.filter_map(|v| v["params"]["level"].as_str().map(str::to_owned))
.collect()
}
#[test]
fn delivers_messages_at_or_above_requested_level() {
let lines = emit_within_request(Some(LoggingLevel::Warning));
let got = levels(&lines);
assert!(got.contains(&"error".to_owned()), "got: {got:?}");
assert!(got.contains(&"warning".to_owned()), "got: {got:?}");
assert!(!got.contains(&"info".to_owned()), "got: {got:?}");
assert!(!got.contains(&"debug".to_owned()), "got: {got:?}");
}
#[test]
fn delivers_everything_at_debug() {
let got = levels(&emit_within_request(Some(LoggingLevel::Debug)));
for lvl in ["error", "warning", "info", "debug"] {
assert!(got.contains(&lvl.to_owned()), "missing {lvl}, got: {got:?}");
}
}
#[test]
fn suppresses_all_messages_without_requested_level() {
let got = levels(&emit_within_request(None));
assert!(got.is_empty(), "expected none, got: {got:?}");
}
#[test]
fn preserves_mcp_specific_severity_past_tracing() {
use crate::types::notification::LogMessage;
let buf = Arc::new(Mutex::new(Vec::new()));
let subscriber = tracing_subscriber::registry()
.with(super::span_context())
.with(
tracing_subscriber::fmt::layer()
.event_format(NotificationFormatter)
.with_writer(BufWriter(buf.clone())),
);
tracing::subscriber::with_default(subscriber, || {
let span = tracing::info_span!(
"request",
mcp_log_level = u64::from(LoggingLevel::Emergency.severity())
);
let _entered = span.enter();
LogMessage::new(LoggingLevel::Emergency, None, None).write();
LogMessage::new(LoggingLevel::Error, None, None).write();
});
let raw = buf.lock().unwrap().clone();
let got: Vec<String> = String::from_utf8(raw)
.unwrap()
.lines()
.filter(|l| !l.trim().is_empty())
.map(str::to_owned)
.collect();
let got = levels(&got);
assert!(got.contains(&"emergency".to_owned()), "got: {got:?}");
assert!(!got.contains(&"error".to_owned()), "got: {got:?}");
}
#[tokio::test]
async fn routes_events_from_nested_spans_to_the_request_sink() {
let session_id = uuid::Uuid::new_v4();
let mut sink_rx = super::super::sink::register(session_id, 8, false).await;
let (fallback_tx, mut fallback_rx) = tokio::sync::mpsc::channel(8);
super::LOG_REGISTRY.register(session_id, 1, fallback_tx);
let subscriber = tracing_subscriber::registry().with(super::MpscLayer);
tracing::subscriber::with_default(subscriber, || {
let request = tracing::info_span!(
"request",
mcp_session_id = session_id.to_string(),
mcp_log_level = u64::from(LoggingLevel::Debug.severity())
);
let _entered = request.enter();
let handler = tracing::info_span!("handler");
let _handler = handler.enter();
tracing::warn!(logger = "tool", "nested message");
});
super::super::sink::unregister(&session_id);
super::LOG_REGISTRY.unregister(&session_id);
let msg = sink_rx
.try_recv()
.expect("an event from a nested span must still reach the request sink");
let json = serde_json::to_value(&msg).unwrap();
assert_eq!(json["method"], "notifications/message");
assert_eq!(json["params"]["data"]["message"], "nested message");
assert!(
fallback_rx.try_recv().is_err(),
"request-scoped notification leaked to the legacy path"
);
}
}