mirror of
https://github.com/openai/codex.git
synced 2026-09-13 11:47:17 +00:00
Propagate RMCP trace context
Add traceparent/tracestate propagation for RMCP streamable HTTP and stdio calls, and cover it with transport-level tests. Co-authored-by: Codex <noreply@openai.com>
This commit is contained in:
5
codex-rs/Cargo.lock
generated
5
codex-rs/Cargo.lock
generated
@@ -2514,6 +2514,7 @@ dependencies = [
|
||||
"axum",
|
||||
"codex-client",
|
||||
"codex-keyring-store",
|
||||
"codex-otel",
|
||||
"codex-protocol",
|
||||
"codex-utils-cargo-bin",
|
||||
"codex-utils-home-dir",
|
||||
@@ -2521,6 +2522,8 @@ dependencies = [
|
||||
"futures",
|
||||
"keyring",
|
||||
"oauth2",
|
||||
"opentelemetry",
|
||||
"opentelemetry_sdk",
|
||||
"pretty_assertions",
|
||||
"reqwest",
|
||||
"rmcp",
|
||||
@@ -2535,6 +2538,8 @@ dependencies = [
|
||||
"tiny_http",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"tracing-opentelemetry",
|
||||
"tracing-subscriber",
|
||||
"urlencoding",
|
||||
"webbrowser",
|
||||
"which 8.0.0",
|
||||
|
||||
@@ -36,6 +36,7 @@ use core_test_support::responses::mount_sse_once;
|
||||
use core_test_support::skip_if_no_network;
|
||||
use core_test_support::stdio_server_bin;
|
||||
use core_test_support::test_codex::test_codex;
|
||||
use core_test_support::tracing::install_test_tracing;
|
||||
use core_test_support::wait_for_event;
|
||||
use core_test_support::wait_for_event_with_timeout;
|
||||
use reqwest::Client;
|
||||
@@ -687,6 +688,8 @@ async fn stdio_server_propagates_whitelisted_env_vars() -> anyhow::Result<()> {
|
||||
async fn streamable_http_tool_call_round_trip() -> anyhow::Result<()> {
|
||||
skip_if_no_network!(Ok(()));
|
||||
|
||||
let _trace_test_context = install_test_tracing("rmcp-integration-tests");
|
||||
|
||||
let server = responses::start_mock_server().await;
|
||||
|
||||
let call_id = "call-456";
|
||||
@@ -733,6 +736,7 @@ async fn streamable_http_tool_call_round_trip() -> anyhow::Result<()> {
|
||||
.kill_on_drop(true)
|
||||
.env("MCP_STREAMABLE_HTTP_BIND_ADDR", &bind_addr)
|
||||
.env("MCP_TEST_VALUE", expected_env_value)
|
||||
.env("MCP_EXPECT_TRACEPARENT", "1")
|
||||
.spawn()?;
|
||||
|
||||
wait_for_streamable_http_server(&mut http_server_child, &bind_addr, Duration::from_secs(5))
|
||||
|
||||
@@ -15,6 +15,7 @@ axum = { workspace = true, default-features = false, features = [
|
||||
] }
|
||||
codex-client = { workspace = true }
|
||||
codex-keyring-store = { workspace = true }
|
||||
codex-otel = { workspace = true }
|
||||
codex-protocol = { workspace = true }
|
||||
codex-utils-pty = { workspace = true }
|
||||
codex-utils-home-dir = { workspace = true }
|
||||
@@ -60,9 +61,13 @@ which = { workspace = true }
|
||||
|
||||
[dev-dependencies]
|
||||
codex-utils-cargo-bin = { workspace = true }
|
||||
opentelemetry = { workspace = true }
|
||||
opentelemetry_sdk = { workspace = true }
|
||||
pretty_assertions = { workspace = true }
|
||||
serial_test = { workspace = true }
|
||||
tempfile = { workspace = true }
|
||||
tracing-opentelemetry = { workspace = true }
|
||||
tracing-subscriber = { workspace = true }
|
||||
[target.'cfg(target_os = "linux")'.dependencies]
|
||||
keyring = { workspace = true, features = ["linux-native-async-persistent"] }
|
||||
|
||||
|
||||
@@ -315,7 +315,7 @@ impl ServerHandler for TestToolServer {
|
||||
async fn call_tool(
|
||||
&self,
|
||||
request: CallToolRequestParams,
|
||||
_context: rmcp::service::RequestContext<rmcp::service::RoleServer>,
|
||||
context: rmcp::service::RequestContext<rmcp::service::RoleServer>,
|
||||
) -> Result<CallToolResult, McpError> {
|
||||
match request.name.as_ref() {
|
||||
"echo" | "echo-tool" => {
|
||||
@@ -333,9 +333,19 @@ impl ServerHandler for TestToolServer {
|
||||
};
|
||||
|
||||
let env_snapshot: HashMap<String, String> = std::env::vars().collect();
|
||||
let traceparent = context
|
||||
.meta
|
||||
.get("x-codex-traceparent")
|
||||
.and_then(serde_json::Value::as_str);
|
||||
let tracestate = context
|
||||
.meta
|
||||
.get("x-codex-tracestate")
|
||||
.and_then(serde_json::Value::as_str);
|
||||
let structured_content = json!({
|
||||
"echo": format!("ECHOING: {}", args.message),
|
||||
"env": env_snapshot.get("MCP_TEST_VALUE"),
|
||||
"traceparent": traceparent,
|
||||
"tracestate": tracestate,
|
||||
});
|
||||
|
||||
Ok(CallToolResult {
|
||||
|
||||
@@ -58,6 +58,7 @@ struct TestToolServer {
|
||||
const MEMO_URI: &str = "memo://codex/example-note";
|
||||
const MEMO_CONTENT: &str = "This is a sample MCP resource served by the rmcp test server.";
|
||||
const MCP_SESSION_ID_HEADER: &str = "mcp-session-id";
|
||||
const TRACEPARENT_HEADER: &str = "traceparent";
|
||||
const SESSION_POST_FAILURE_CONTROL_PATH: &str = "/test/control/session-post-failure";
|
||||
|
||||
impl TestToolServer {
|
||||
@@ -347,6 +348,11 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
} else {
|
||||
router
|
||||
};
|
||||
let router = if std::env::var("MCP_EXPECT_TRACEPARENT").is_ok() {
|
||||
router.layer(middleware::from_fn(require_traceparent_on_session_post))
|
||||
} else {
|
||||
router
|
||||
};
|
||||
|
||||
axum::serve(listener, router).await?;
|
||||
task::yield_now().await;
|
||||
@@ -389,6 +395,23 @@ async fn arm_session_post_failure(
|
||||
Ok(StatusCode::NO_CONTENT)
|
||||
}
|
||||
|
||||
async fn require_traceparent_on_session_post(request: Request<Body>, next: Next) -> Response {
|
||||
if request.uri().path() != "/mcp"
|
||||
|| request.method() != Method::POST
|
||||
|| !request.headers().contains_key(MCP_SESSION_ID_HEADER)
|
||||
{
|
||||
return next.run(request).await;
|
||||
}
|
||||
|
||||
if request.headers().contains_key(TRACEPARENT_HEADER) {
|
||||
next.run(request).await
|
||||
} else {
|
||||
let mut response = Response::new(Body::from("missing traceparent header"));
|
||||
*response.status_mut() = StatusCode::BAD_REQUEST;
|
||||
response
|
||||
}
|
||||
}
|
||||
|
||||
async fn fail_session_post_when_armed(
|
||||
State(state): State<SessionFailureState>,
|
||||
request: Request<Body>,
|
||||
|
||||
@@ -10,6 +10,9 @@ use std::time::Duration;
|
||||
use anyhow::Result;
|
||||
use anyhow::anyhow;
|
||||
use codex_client::build_reqwest_client_with_custom_ca;
|
||||
use codex_otel::current_span_w3c_trace_context;
|
||||
use codex_otel::span_w3c_trace_context;
|
||||
use codex_protocol::protocol::W3cTraceContext;
|
||||
use futures::FutureExt;
|
||||
use futures::StreamExt;
|
||||
use futures::future::BoxFuture;
|
||||
@@ -64,6 +67,8 @@ use tokio::io::BufReader;
|
||||
use tokio::process::Command;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::time;
|
||||
use tracing::Instrument;
|
||||
use tracing::field::Empty;
|
||||
use tracing::info;
|
||||
use tracing::warn;
|
||||
|
||||
@@ -82,6 +87,14 @@ const JSON_MIME_TYPE: &str = "application/json";
|
||||
const HEADER_LAST_EVENT_ID: &str = "Last-Event-Id";
|
||||
const HEADER_SESSION_ID: &str = "Mcp-Session-Id";
|
||||
const NON_JSON_RESPONSE_BODY_PREVIEW_BYTES: usize = 8_192;
|
||||
const TRACEPARENT_HEADER: &str = "traceparent";
|
||||
const TRACESTATE_HEADER: &str = "tracestate";
|
||||
const TRACEPARENT_META_KEY: &str = "x-codex-traceparent";
|
||||
const TRACESTATE_META_KEY: &str = "x-codex-tracestate";
|
||||
|
||||
tokio::task_local! {
|
||||
static SERVICE_OPERATION_TRACE_CONTEXT: Option<W3cTraceContext>;
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct StreamableHttpResponseClient {
|
||||
@@ -98,6 +111,24 @@ impl StreamableHttpResponseClient {
|
||||
) -> StreamableHttpError<StreamableHttpResponseClientError> {
|
||||
StreamableHttpError::Client(StreamableHttpResponseClientError::from(error))
|
||||
}
|
||||
|
||||
fn apply_trace_context(
|
||||
request: reqwest::RequestBuilder,
|
||||
trace: Option<&W3cTraceContext>,
|
||||
) -> reqwest::RequestBuilder {
|
||||
let Some(trace) = trace else {
|
||||
return request;
|
||||
};
|
||||
|
||||
let mut request = request;
|
||||
if let Some(traceparent) = trace.traceparent.as_deref() {
|
||||
request = request.header(TRACEPARENT_HEADER, traceparent);
|
||||
}
|
||||
if let Some(tracestate) = trace.tracestate.as_deref() {
|
||||
request = request.header(TRACESTATE_HEADER, tracestate);
|
||||
}
|
||||
request
|
||||
}
|
||||
}
|
||||
|
||||
fn build_http_client(default_headers: &HeaderMap) -> Result<reqwest::Client> {
|
||||
@@ -123,6 +154,7 @@ impl StreamableHttpClient for StreamableHttpResponseClient {
|
||||
session_id: Option<Arc<str>>,
|
||||
auth_token: Option<String>,
|
||||
) -> std::result::Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
|
||||
let trace = current_service_operation_trace_context();
|
||||
let mut request = self
|
||||
.inner
|
||||
.post(uri.as_ref())
|
||||
@@ -133,6 +165,7 @@ impl StreamableHttpClient for StreamableHttpResponseClient {
|
||||
if let Some(session_id_value) = session_id.as_ref() {
|
||||
request = request.header(HEADER_SESSION_ID, session_id_value.as_ref());
|
||||
}
|
||||
request = Self::apply_trace_context(request, trace.as_ref());
|
||||
|
||||
let response = request
|
||||
.json(&message)
|
||||
@@ -224,10 +257,12 @@ impl StreamableHttpClient for StreamableHttpResponseClient {
|
||||
session: Arc<str>,
|
||||
auth_token: Option<String>,
|
||||
) -> std::result::Result<(), StreamableHttpError<Self::Error>> {
|
||||
let trace = current_service_operation_trace_context();
|
||||
let mut request_builder = self.inner.delete(uri.as_ref());
|
||||
if let Some(auth_header) = auth_token {
|
||||
request_builder = request_builder.bearer_auth(auth_header);
|
||||
}
|
||||
request_builder = Self::apply_trace_context(request_builder, trace.as_ref());
|
||||
let response = request_builder
|
||||
.header(HEADER_SESSION_ID, session.as_ref())
|
||||
.send()
|
||||
@@ -254,6 +289,7 @@ impl StreamableHttpClient for StreamableHttpResponseClient {
|
||||
BoxStream<'static, std::result::Result<Sse, sse_stream::Error>>,
|
||||
StreamableHttpError<Self::Error>,
|
||||
> {
|
||||
let trace = current_service_operation_trace_context();
|
||||
let mut request_builder = self
|
||||
.inner
|
||||
.get(uri.as_ref())
|
||||
@@ -265,6 +301,7 @@ impl StreamableHttpClient for StreamableHttpResponseClient {
|
||||
if let Some(auth_header) = auth_token {
|
||||
request_builder = request_builder.bearer_auth(auth_header);
|
||||
}
|
||||
request_builder = Self::apply_trace_context(request_builder, trace.as_ref());
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
@@ -733,6 +770,8 @@ impl RmcpClient {
|
||||
let rmcp_params = rmcp_params.clone();
|
||||
let meta = meta.clone();
|
||||
async move {
|
||||
let trace = current_service_operation_trace_context();
|
||||
let meta = merge_trace_context_into_meta(meta, trace.as_ref());
|
||||
let result = service
|
||||
.peer()
|
||||
.send_request_with_option(
|
||||
@@ -1052,41 +1091,104 @@ impl RmcpClient {
|
||||
Fut: std::future::Future<Output = std::result::Result<T, rmcp::service::ServiceError>>,
|
||||
{
|
||||
let service = self.service().await?;
|
||||
match Self::run_service_operation_once(Arc::clone(&service), label, timeout, &operation)
|
||||
.await
|
||||
let operation_span = self.service_operation_span(label);
|
||||
let operation_trace = span_w3c_trace_context(&operation_span);
|
||||
match Self::run_service_operation_once(
|
||||
Arc::clone(&service),
|
||||
label,
|
||||
timeout,
|
||||
operation_span,
|
||||
operation_trace,
|
||||
&operation,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => Ok(result),
|
||||
Err(error) if Self::is_session_expired_404(&error) => {
|
||||
self.reinitialize_after_session_expiry(&service).await?;
|
||||
let recovered_service = self.service().await?;
|
||||
Self::run_service_operation_once(recovered_service, label, timeout, &operation)
|
||||
.await
|
||||
.map_err(Into::into)
|
||||
let operation_span = self.service_operation_span(label);
|
||||
let operation_trace = span_w3c_trace_context(&operation_span);
|
||||
Self::run_service_operation_once(
|
||||
recovered_service,
|
||||
label,
|
||||
timeout,
|
||||
operation_span,
|
||||
operation_trace,
|
||||
&operation,
|
||||
)
|
||||
.await
|
||||
.map_err(Into::into)
|
||||
}
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
fn service_operation_span(&self, label: &str) -> tracing::Span {
|
||||
let span = tracing::info_span!(
|
||||
"mcp.client.operation",
|
||||
otel.kind = "client",
|
||||
rpc.system = "jsonrpc",
|
||||
rpc.method = label,
|
||||
mcp.transport = Empty,
|
||||
mcp.server.name = Empty,
|
||||
server.address = Empty,
|
||||
server.port = Empty,
|
||||
);
|
||||
|
||||
match &self.transport_recipe {
|
||||
TransportRecipe::Stdio { .. } => {
|
||||
span.record("mcp.transport", "stdio");
|
||||
}
|
||||
TransportRecipe::StreamableHttp {
|
||||
server_name, url, ..
|
||||
} => {
|
||||
span.record("mcp.transport", "streamable_http");
|
||||
span.record("mcp.server.name", server_name.as_str());
|
||||
if let Ok(parsed_url) = reqwest::Url::parse(url) {
|
||||
if let Some(host) = parsed_url.host_str() {
|
||||
span.record("server.address", host);
|
||||
}
|
||||
if let Some(port) = parsed_url.port_or_known_default() {
|
||||
span.record("server.port", port as i64);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
span
|
||||
}
|
||||
|
||||
async fn run_service_operation_once<T, F, Fut>(
|
||||
service: Arc<RunningService<RoleClient, LoggingClientHandler>>,
|
||||
label: &str,
|
||||
timeout: Option<Duration>,
|
||||
operation_span: tracing::Span,
|
||||
operation_trace: Option<W3cTraceContext>,
|
||||
operation: &F,
|
||||
) -> std::result::Result<T, ClientOperationError>
|
||||
where
|
||||
F: Fn(Arc<RunningService<RoleClient, LoggingClientHandler>>) -> Fut,
|
||||
Fut: std::future::Future<Output = std::result::Result<T, rmcp::service::ServiceError>>,
|
||||
{
|
||||
match timeout {
|
||||
Some(duration) => time::timeout(duration, operation(service))
|
||||
SERVICE_OPERATION_TRACE_CONTEXT
|
||||
.scope(operation_trace, async move {
|
||||
async move {
|
||||
match timeout {
|
||||
Some(duration) => time::timeout(duration, operation(service))
|
||||
.await
|
||||
.map_err(|_| ClientOperationError::Timeout {
|
||||
label: label.to_string(),
|
||||
duration,
|
||||
})?
|
||||
.map_err(ClientOperationError::from),
|
||||
None => operation(service).await.map_err(ClientOperationError::from),
|
||||
}
|
||||
}
|
||||
.instrument(operation_span)
|
||||
.await
|
||||
.map_err(|_| ClientOperationError::Timeout {
|
||||
label: label.to_string(),
|
||||
duration,
|
||||
})?
|
||||
.map_err(ClientOperationError::from),
|
||||
None => operation(service).await.map_err(ClientOperationError::from),
|
||||
}
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
fn is_session_expired_404(error: &ClientOperationError) -> bool {
|
||||
@@ -1161,6 +1263,39 @@ impl RmcpClient {
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_trace_context_into_meta(
|
||||
meta: Option<rmcp::model::Meta>,
|
||||
trace: Option<&W3cTraceContext>,
|
||||
) -> Option<rmcp::model::Meta> {
|
||||
let mut meta = meta.unwrap_or_default();
|
||||
let Some(trace) = trace else {
|
||||
return (!meta.is_empty()).then_some(meta);
|
||||
};
|
||||
|
||||
if let Some(traceparent) = trace.traceparent.as_ref() {
|
||||
meta.insert(
|
||||
TRACEPARENT_META_KEY.to_string(),
|
||||
serde_json::Value::String(traceparent.clone()),
|
||||
);
|
||||
}
|
||||
if let Some(tracestate) = trace.tracestate.as_ref() {
|
||||
meta.insert(
|
||||
TRACESTATE_META_KEY.to_string(),
|
||||
serde_json::Value::String(tracestate.clone()),
|
||||
);
|
||||
}
|
||||
|
||||
(!meta.is_empty()).then_some(meta)
|
||||
}
|
||||
|
||||
fn current_service_operation_trace_context() -> Option<W3cTraceContext> {
|
||||
SERVICE_OPERATION_TRACE_CONTEXT
|
||||
.try_with(Clone::clone)
|
||||
.ok()
|
||||
.flatten()
|
||||
.or_else(current_span_w3c_trace_context)
|
||||
}
|
||||
|
||||
async fn create_oauth_transport_and_runtime(
|
||||
server_name: &str,
|
||||
url: &str,
|
||||
@@ -1207,3 +1342,75 @@ async fn create_oauth_transport_and_runtime(
|
||||
|
||||
Ok((transport, runtime))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod trace_tests {
|
||||
use super::StreamableHttpResponseClient;
|
||||
use super::TRACEPARENT_HEADER;
|
||||
use super::TRACEPARENT_META_KEY;
|
||||
use super::TRACESTATE_HEADER;
|
||||
use super::TRACESTATE_META_KEY;
|
||||
use super::merge_trace_context_into_meta;
|
||||
use codex_protocol::protocol::W3cTraceContext;
|
||||
use pretty_assertions::assert_eq;
|
||||
use reqwest::Client;
|
||||
use serde_json::Value;
|
||||
|
||||
#[test]
|
||||
fn merge_trace_context_into_meta_preserves_existing_fields() {
|
||||
let trace = W3cTraceContext {
|
||||
traceparent: Some("00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01".into()),
|
||||
tracestate: Some("vendor=value".into()),
|
||||
};
|
||||
let meta = rmcp::model::Meta(serde_json::Map::from_iter([(
|
||||
"existing".to_string(),
|
||||
Value::String("value".into()),
|
||||
)]));
|
||||
|
||||
let merged = merge_trace_context_into_meta(Some(meta), Some(&trace)).expect("meta");
|
||||
|
||||
assert_eq!(
|
||||
merged,
|
||||
rmcp::model::Meta(serde_json::Map::from_iter([
|
||||
("existing".to_string(), Value::String("value".into())),
|
||||
(
|
||||
TRACEPARENT_META_KEY.to_string(),
|
||||
Value::String("00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01".into())
|
||||
),
|
||||
(
|
||||
TRACESTATE_META_KEY.to_string(),
|
||||
Value::String("vendor=value".into())
|
||||
),
|
||||
]))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn apply_trace_context_injects_http_headers() {
|
||||
let trace = W3cTraceContext {
|
||||
traceparent: Some("00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01".into()),
|
||||
tracestate: Some("vendor=value".into()),
|
||||
};
|
||||
let request = StreamableHttpResponseClient::apply_trace_context(
|
||||
Client::new().post("http://example.com"),
|
||||
Some(&trace),
|
||||
)
|
||||
.build()
|
||||
.expect("request");
|
||||
|
||||
assert_eq!(
|
||||
request
|
||||
.headers()
|
||||
.get(TRACEPARENT_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01")
|
||||
);
|
||||
assert_eq!(
|
||||
request
|
||||
.headers()
|
||||
.get(TRACESTATE_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("vendor=value")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,10 @@ use codex_rmcp_client::ElicitationResponse;
|
||||
use codex_rmcp_client::RmcpClient;
|
||||
use codex_utils_cargo_bin::CargoBinError;
|
||||
use futures::FutureExt as _;
|
||||
use opentelemetry::global;
|
||||
use opentelemetry::trace::TracerProvider as _;
|
||||
use opentelemetry_sdk::propagation::TraceContextPropagator;
|
||||
use opentelemetry_sdk::trace::SdkTracerProvider;
|
||||
use rmcp::model::AnnotateAble;
|
||||
use rmcp::model::ClientCapabilities;
|
||||
use rmcp::model::ElicitationCapability;
|
||||
@@ -18,6 +22,10 @@ use rmcp::model::ProtocolVersion;
|
||||
use rmcp::model::ReadResourceRequestParams;
|
||||
use rmcp::model::ResourceContents;
|
||||
use serde_json::json;
|
||||
use tracing::Instrument;
|
||||
use tracing::dispatcher::DefaultGuard;
|
||||
use tracing_subscriber::layer::SubscriberExt;
|
||||
use tracing_subscriber::util::SubscriberInitExt;
|
||||
|
||||
const RESOURCE_URI: &str = "memo://codex/example-note";
|
||||
|
||||
@@ -53,6 +61,25 @@ fn init_params() -> InitializeRequestParams {
|
||||
}
|
||||
}
|
||||
|
||||
struct TestTracingContext {
|
||||
_provider: SdkTracerProvider,
|
||||
_guard: DefaultGuard,
|
||||
}
|
||||
|
||||
fn install_test_tracing(tracer_name: &str) -> TestTracingContext {
|
||||
global::set_text_map_propagator(TraceContextPropagator::new());
|
||||
|
||||
let provider = SdkTracerProvider::builder().build();
|
||||
let tracer = provider.tracer(tracer_name.to_string());
|
||||
let subscriber =
|
||||
tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer));
|
||||
|
||||
TestTracingContext {
|
||||
_provider: provider,
|
||||
_guard: subscriber.set_default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
|
||||
async fn rmcp_client_can_list_and_read_resources() -> anyhow::Result<()> {
|
||||
let client = RmcpClient::new_stdio_client(
|
||||
@@ -149,3 +176,61 @@ async fn rmcp_client_can_list_and_read_resources() -> anyhow::Result<()> {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn stdio_tool_call_propagates_trace_metadata() -> anyhow::Result<()> {
|
||||
let _trace = install_test_tracing("rmcp-stdio-trace-test");
|
||||
let client = RmcpClient::new_stdio_client(
|
||||
stdio_server_bin()?.into(),
|
||||
Vec::<OsString>::new(),
|
||||
None,
|
||||
&[],
|
||||
None,
|
||||
)
|
||||
.await?;
|
||||
|
||||
client
|
||||
.initialize(
|
||||
init_params(),
|
||||
Some(Duration::from_secs(5)),
|
||||
Box::new(|_, _| {
|
||||
async {
|
||||
Ok(ElicitationResponse {
|
||||
action: ElicitationAction::Accept,
|
||||
content: Some(json!({})),
|
||||
meta: None,
|
||||
})
|
||||
}
|
||||
.boxed()
|
||||
}),
|
||||
)
|
||||
.await?;
|
||||
|
||||
let result = async {
|
||||
client
|
||||
.call_tool(
|
||||
"echo".to_string(),
|
||||
Some(json!({ "message": "ping" })),
|
||||
None,
|
||||
Some(Duration::from_secs(5)),
|
||||
)
|
||||
.await
|
||||
}
|
||||
.instrument(tracing::info_span!("rmcp.client.trace_test"))
|
||||
.await?;
|
||||
|
||||
assert_eq!(result.is_error, Some(false));
|
||||
let structured = result.structured_content.expect("structured content");
|
||||
assert_eq!(structured["echo"], json!("ECHOING: ping"));
|
||||
assert_eq!(structured["env"], serde_json::Value::Null);
|
||||
assert!(
|
||||
structured["tracestate"].is_null()
|
||||
|| structured["tracestate"].as_str().is_some_and(str::is_empty)
|
||||
);
|
||||
let traceparent = structured["traceparent"]
|
||||
.as_str()
|
||||
.expect("traceparent should be propagated via request metadata");
|
||||
assert!(traceparent.starts_with("00-"));
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user