Skip to content
Merged
Show file tree
Hide file tree
Changes from 9 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions .changesets/feat_david_support_traceparent_context.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
### Server adds support for incoming distributed trace context propagation
Comment thread
david-castaneda marked this conversation as resolved.
Outdated

The MCP server now extracts W3C traceparent headers from incoming requests and uses this context for its own emitted traces, enabling handler spans to nest under parent traces for complete end-to-end observability.
1 change: 1 addition & 0 deletions crates/apollo-mcp-server/src/server/states.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ mod operations_configured;
mod running;
mod schema_configured;
mod starting;
mod telemetry;

use configuring::Configuring;
use operations_configured::OperationsConfigured;
Expand Down
9 changes: 5 additions & 4 deletions crates/apollo-mcp-server/src/server/states/running.rs
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ use crate::apps::find_and_execute_app;
use crate::generated::telemetry::{TelemetryAttribute, TelemetryMetric};
use crate::meter;
use crate::operations::{execute_operation, find_and_execute_operation};
use crate::server::states::telemetry::get_parent_span;
use crate::{
apps::AppResource,
custom_scalar_map::CustomScalarMap,
Expand Down Expand Up @@ -269,7 +270,7 @@ impl Running {
}

impl ServerHandler for Running {
#[tracing::instrument(skip_all, fields(apollo.mcp.client_name = request.client_info.name, apollo.mcp.client_version = request.client_info.version))]
#[tracing::instrument(skip_all, parent = get_parent_span(&context), fields(apollo.mcp.client_name = request.client_info.name, apollo.mcp.client_version = request.client_info.version))]
async fn initialize(
&self,
request: InitializeRequestParam,
Expand All @@ -296,7 +297,7 @@ impl ServerHandler for Running {
Ok(self.get_info())
}

#[tracing::instrument(skip_all, fields(apollo.mcp.tool_name = request.name.as_ref(), apollo.mcp.request_id = %context.id.clone()))]
#[tracing::instrument(skip_all, parent = get_parent_span(&context), fields(apollo.mcp.tool_name = request.name.as_ref(), apollo.mcp.request_id = %context.id.clone()))]
async fn call_tool(
&self,
request: CallToolRequestParam,
Expand Down Expand Up @@ -408,11 +409,11 @@ impl ServerHandler for Running {
result
}

#[tracing::instrument(skip_all)]
#[tracing::instrument(skip_all, parent = get_parent_span(&context))]
async fn list_tools(
&self,
_request: Option<PaginatedRequestParam>,
_context: RequestContext<RoleServer>,
context: RequestContext<RoleServer>,
) -> Result<ListToolsResult, McpError> {
let meter = &meter::METER;
meter
Expand Down
88 changes: 55 additions & 33 deletions crates/apollo-mcp-server/src/server/states/starting.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,9 @@ use rmcp::{
use serde_json::json;
use tokio::sync::RwLock;
use tokio_util::sync::CancellationToken;
use tower_http::trace::TraceLayer;
use tracing::{Instrument as _, debug, error, info, trace};

use crate::server::states::telemetry::otel_context_middleware;
use crate::{
errors::ServerError,
explorer::Explorer,
Expand Down Expand Up @@ -228,38 +228,9 @@ impl Starting {
.layer(HttpMetricsLayerBuilder::new().build())
// include trace context as header into the response
.layer(OtelInResponseLayer)
//start OpenTelemetry trace on incoming request
.layer(OtelAxumLayer::default())
// Add tower-http tracing layer for additional HTTP-level tracing
.layer(
TraceLayer::new_for_http()
.make_span_with(|request: &axum::http::Request<_>| {
tracing::info_span!(
"mcp_server",
method = %request.method(),
uri = %request.uri(),
session_id = tracing::field::Empty,
status_code = tracing::field::Empty,
)
})
.on_response(
|response: &axum::http::Response<_>,
_latency: std::time::Duration,
span: &tracing::Span| {
span.record(
"status_code",
tracing::field::display(response.status()),
);
if let Some(session_id) = response
.headers()
.get("mcp-session-id")
.and_then(|v| v.to_str().ok())
{
span.record("session_id", tracing::field::display(session_id));
}
},
),
);
// start OpenTelemetry trace on incoming request
.layer(axum::middleware::from_fn(otel_context_middleware))
.layer(OtelAxumLayer::default());

// Add health check endpoint if configured
if let Some(health_check) = health_check.filter(|h| h.config().enabled) {
Expand Down Expand Up @@ -361,7 +332,10 @@ async fn health_endpoint(

#[cfg(test)]
mod tests {
use axum::{body::Body, http::Request};
use http::HeaderMap;
use tower::ServiceExt;
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
use url::Url;

use crate::health::HealthCheckConfig;
Expand Down Expand Up @@ -408,4 +382,52 @@ mod tests {
let running = starting.start();
assert!(running.await.is_ok());
}

#[tokio::test]
async fn test_otel_context_middleware_does_not_break_requests() {
use tracing_subscriber::layer::SubscriberExt;

// Initialize tracing for test
let _ = tracing_subscriber::registry()
.with(tracing_subscriber::fmt::layer())
.try_init();

// Create a simple router with the middleware
let app = axum::Router::new()
.route("/test", axum::routing::get(|| async { "ok" }))
.layer(axum::middleware::from_fn(otel_context_middleware));

// Valid W3C traceparent header
let trace_id = "4bf92f3577b34da6a3ce929d0e0e4736";
let span_id = "00f067aa0ba902b7";
let traceparent = format!("00-{}-{}-01", trace_id, span_id);

let request = Request::builder()
.uri("/test")
.header("traceparent", traceparent)
.body(Body::empty())
.unwrap();

let response = app.oneshot(request).await.unwrap();

// If we got a response, the middleware worked
assert_eq!(response.status(), 200);
}

#[tokio::test]
async fn test_otel_context_middleware_works_without_traceparent() {
let _ = tracing_subscriber::registry()
.with(tracing_subscriber::fmt::layer())
.try_init();

let app = axum::Router::new()
.route("/test", axum::routing::get(|| async { "ok" }))
.layer(axum::middleware::from_fn(otel_context_middleware));

let request = Request::builder().uri("/test").body(Body::empty()).unwrap();

let response = app.oneshot(request).await.unwrap();

assert_eq!(response.status(), 200);
}
}
65 changes: 65 additions & 0 deletions crates/apollo-mcp-server/src/server/states/telemetry.rs
Comment thread
david-castaneda marked this conversation as resolved.
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
use axum::extract::Request;
use axum::middleware::Next;
use axum::response::Response;
use opentelemetry::global;
use opentelemetry::propagation::Extractor;
use rmcp::RoleServer;
use rmcp::service::RequestContext;
use tracing::Instrument;
use tracing_opentelemetry::OpenTelemetrySpanExt;

// Custom extractor for axum headers
struct HeaderExtractor<'a>(&'a axum::http::HeaderMap);

// Implement the Extractor trait for HeaderExtractor
impl<'a> Extractor for HeaderExtractor<'a> {
fn get(&self, key: &str) -> Option<&str> {
self.0.get(key).and_then(|v| v.to_str().ok())
}

fn keys(&self) -> Vec<&str> {
self.0.keys().map(|k| k.as_str()).collect()
}
}

/// Middleware that extracts and stores OpenTelemetry context in request extensions
pub async fn otel_context_middleware(mut request: Request, next: Next) -> Response {
let parent_cx = global::get_text_map_propagator(|propagator| {
propagator.extract(&HeaderExtractor(request.headers()))
});

let span = tracing::info_span!(
"mcp_server",
method = %request.method(),
uri = %request.uri(),
session_id = tracing::field::Empty,
status_code = tracing::field::Empty,
);
span.set_parent(parent_cx);

request.extensions_mut().insert(span.clone()); // Store the span in request extensions

let response = next.run(request).instrument(span.clone()).await;

span.record("status_code", tracing::field::display(response.status()));

if let Some(session_id) = response
.headers()
.get("mcp-session-id")
.and_then(|v| v.to_str().ok())
{
span.record("session_id", tracing::field::display(session_id));
}

response
}

// Helper function to retrieve the parent span from the request context
pub fn get_parent_span(context: &RequestContext<RoleServer>) -> tracing::Span {
context
.extensions
.get::<axum::http::request::Parts>()
.and_then(|parts| parts.extensions.get::<tracing::Span>())
.cloned()
.unwrap_or_else(tracing::Span::current)
}