use schemars::JsonSchema;
use serde::Deserialize;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::time::{Duration, timeout};
#[cfg(feature = "stateless")]
use tower_mcp::ProtocolSupport;
#[cfg(feature = "stateless")]
use tower_mcp::extract::RawArgs;
use tower_mcp::extract::{Context, Json};
use tower_mcp::protocol::{ElicitAction, ElicitFormParams, ElicitFormSchema};
use tower_mcp::transport::stdio::BidirectionalStdioTransport;
use tower_mcp::{CallToolResult, McpRouter, StdioTransport, ToolBuilder};
use tower_mcp_types::testing::assert_jsonrpc_error_response;
fn router() -> McpRouter {
McpRouter::new().server_info("stdio-loop-test", "0.0.0")
}
async fn read_n_frames<R>(mut reader: BufReader<R>, expected: usize) -> Vec<serde_json::Value>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut out = Vec::with_capacity(expected);
while out.len() < expected {
let mut line = String::new();
let n = reader
.read_line(&mut line)
.await
.expect("read from server output");
if n == 0 {
break; }
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let v: serde_json::Value = serde_json::from_str(trimmed)
.unwrap_or_else(|e| panic!("invalid JSON on output: {e}: {trimmed}"));
out.push(v);
}
out
}
mod safe_frame_tracing {
use super::*;
use std::io::Write;
use std::sync::{Arc, Mutex};
use tower_mcp::GenericStdioTransport;
use tower_mcp::context::{ServerNotification, notification_channel};
use tower_mcp::protocol::{LogLevel, LoggingMessageParams};
static TRACE_CAPTURE_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
#[derive(Clone, Default)]
struct CaptureWriter(Arc<Mutex<Vec<u8>>>);
impl Write for CaptureWriter {
fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
self.0
.lock()
.expect("trace capture lock")
.extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
impl CaptureWriter {
fn contents(&self) -> String {
String::from_utf8(self.0.lock().expect("trace capture lock").clone())
.expect("tracing output is UTF-8")
}
}
async fn trace_ping<F, Fut>(sentinel: &str, run: F) -> (usize, usize)
where
F: FnOnce(tokio::io::DuplexStream, tokio::io::DuplexStream) -> Fut,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let frame = format!(r#"{{"jsonrpc":"2.0","id":"{sentinel}","method":"ping"}}"#);
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let transport = tokio::spawn(run(server_stdin, server_stdout));
stdin_writer.write_all(frame.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let frames = timeout(
Duration::from_secs(5),
read_n_frames(BufReader::new(server_stdout_reader), 1),
)
.await
.expect("stdio transport must answer the ping");
timeout(Duration::from_secs(5), transport)
.await
.expect("stdio transport must stop at EOF")
.expect("transport task");
assert_eq!(frames.len(), 1, "expected one ping response: {frames:?}");
assert_eq!(frames[0]["id"], sentinel);
let response_bytes = serde_json::to_string(&frames[0]).unwrap().len();
(frame.len(), response_bytes)
}
#[tokio::test(flavor = "current_thread")]
async fn async_stdio_variants_trace_metadata_without_frame_bodies() {
const PLAIN_SECRET: &str = "plain-stdio-secret-é-e8c2";
const LAYERED_SECRET: &str = "layered-stdio-secret-b50a";
const GENERIC_SECRET: &str = "generic-stdio-secret-7f43";
const BIDI_SECRET: &str = "bidi-stdio-secret-1d69";
let _trace_capture = TRACE_CAPTURE_LOCK.lock().await;
let captured = CaptureWriter::default();
let trace_output = captured.clone();
let tracing = tracing_subscriber::fmt()
.with_ansi(false)
.without_time()
.with_max_level(tracing::Level::DEBUG)
.with_writer(move || captured.clone())
.finish();
let _subscriber = tracing::subscriber::set_default(tracing);
let mut byte_lengths = Vec::new();
byte_lengths.push(
trace_ping(PLAIN_SECRET, |stdin, stdout| async move {
StdioTransport::new(router())
.run_with_streams(stdin, stdout)
.await
.expect("plain stdio loop");
})
.await,
);
byte_lengths.push(
trace_ping(LAYERED_SECRET, |stdin, stdout| async move {
StdioTransport::new(router())
.layer(tower::layer::util::Identity::new())
.run_with_streams(stdin, stdout)
.await
.expect("layered stdio loop");
})
.await,
);
byte_lengths.push(
trace_ping(GENERIC_SECRET, |stdin, stdout| async move {
GenericStdioTransport::new(router())
.run_with_streams(stdin, stdout)
.await
.expect("generic stdio loop");
})
.await,
);
byte_lengths.push(
trace_ping(BIDI_SECRET, |stdin, stdout| async move {
BidirectionalStdioTransport::new(router())
.run_with_streams(stdin, stdout)
.await
.expect("bidirectional stdio loop");
})
.await,
);
let traces = trace_output.contents();
for secret in [PLAIN_SECRET, LAYERED_SECRET, GENERIC_SECRET, BIDI_SECRET] {
assert!(!traces.contains(secret), "frame body leaked: {traces}");
}
assert!(
!traces.contains("input="),
"raw input field returned: {traces}"
);
assert!(
!traces.contains("output="),
"raw output field returned: {traces}"
);
assert_eq!(
traces.matches("direction=\"inbound\"").count(),
4,
"{traces}"
);
assert_eq!(
traces.matches("direction=\"outbound\"").count(),
4,
"{traces}"
);
assert_eq!(
traces.matches("frame_kind=\"message\"").count(),
4,
"{traces}"
);
assert_eq!(
traces.matches("frame_kind=\"response\"").count(),
4,
"{traces}"
);
for (inbound, outbound) in byte_lengths {
assert!(
traces.contains(&format!("utf8_bytes={inbound}")),
"{traces}"
);
assert!(
traces.contains(&format!("utf8_bytes={outbound}")),
"{traces}"
);
}
}
#[tokio::test(flavor = "current_thread")]
async fn bidi_outgoing_request_trace_omits_method_and_params() {
const SECRET_METHOD: &str = "private/provider-request-57d1";
const SECRET_PARAM: &str = "provider-session-token-43d9";
let _trace_capture = TRACE_CAPTURE_LOCK.lock().await;
let captured = CaptureWriter::default();
let trace_output = captured.clone();
let tracing = tracing_subscriber::fmt()
.with_ansi(false)
.without_time()
.with_max_level(tracing::Level::DEBUG)
.with_writer(move || captured.clone())
.finish();
let _subscriber = tracing::subscriber::set_default(tracing);
let mut transport = BidirectionalStdioTransport::new(router());
let requester = transport.client_requester();
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let transport = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
.expect("bidirectional stdio loop");
});
let request = tokio::spawn(async move {
requester
.request(
SECRET_METHOD.to_string(),
serde_json::json!({ "credential": SECRET_PARAM }),
)
.await
});
let mut stdout_reader = BufReader::new(server_stdout_reader);
let outgoing = timeout(Duration::from_secs(5), read_frame(&mut stdout_reader))
.await
.expect("server-to-client request");
assert_eq!(outgoing["method"], SECRET_METHOD);
assert_eq!(outgoing["params"]["credential"], SECRET_PARAM);
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": outgoing["id"],
"result": { "accepted": true }
});
stdin_writer
.write_all(format!("{response}\n").as_bytes())
.await
.unwrap();
stdin_writer.flush().await.unwrap();
timeout(Duration::from_secs(5), request)
.await
.expect("request round trip")
.expect("request task")
.expect("request result");
drop(stdin_writer);
timeout(Duration::from_secs(5), transport)
.await
.expect("transport stops at EOF")
.expect("transport task");
let traces = trace_output.contents();
assert!(!traces.contains(SECRET_METHOD), "method leaked: {traces}");
assert!(!traces.contains(SECRET_PARAM), "params leaked: {traces}");
assert!(traces.contains("direction=\"outbound\""), "{traces}");
assert!(traces.contains("frame_kind=\"request\""), "{traces}");
}
#[tokio::test(flavor = "current_thread")]
async fn generic_notification_trace_omits_notification_body() {
const SECRET: &str = "notification-terminal-output-82ac";
let _trace_capture = TRACE_CAPTURE_LOCK.lock().await;
let captured = CaptureWriter::default();
let trace_output = captured.clone();
let tracing = tracing_subscriber::fmt()
.with_ansi(false)
.without_time()
.with_max_level(tracing::Level::DEBUG)
.with_writer(move || captured.clone())
.finish();
let _subscriber = tracing::subscriber::set_default(tracing);
let (notification_tx, notification_rx) = notification_channel(4);
let mut transport = GenericStdioTransport::with_notifications(router(), notification_rx);
let (stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let transport = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
.expect("generic stdio loop");
});
notification_tx
.send(ServerNotification::LogMessage(LoggingMessageParams {
level: LogLevel::Info,
logger: Some("test".to_string()),
data: serde_json::json!(SECRET),
meta: None,
}))
.await
.expect("notification receiver");
let mut stdout_reader = BufReader::new(server_stdout_reader);
let notification = timeout(Duration::from_secs(5), read_frame(&mut stdout_reader))
.await
.expect("notification frame");
assert_eq!(notification["params"]["data"], SECRET);
drop(stdin_writer);
timeout(Duration::from_secs(5), transport)
.await
.expect("transport stops at EOF")
.expect("transport task");
let traces = trace_output.contents();
assert!(!traces.contains(SECRET), "notification leaked: {traces}");
assert!(traces.contains("direction=\"outbound\""), "{traces}");
assert!(traces.contains("frame_kind=\"notification\""), "{traces}");
}
}
#[tokio::test]
async fn stdio_transport_parse_error_wire_shape() {
let mut transport = StdioTransport::new(router());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut stdin_writer = server_stdin_writer;
stdin_writer
.write_all(b"not valid json{{{\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let reader = BufReader::new(server_stdout_reader);
let frames = read_n_frames(reader, 1).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(
frames.len(),
1,
"expected one parse-error frame, got: {frames:?}"
);
let frame = &frames[0];
assert_jsonrpc_error_response(frame);
assert!(
frame["id"].is_null(),
"parse error id must be null, got: {frame}"
);
assert_eq!(frame["error"]["code"].as_i64().unwrap(), -32700);
}
#[tokio::test]
async fn stdio_transport_loop_continues_after_parse_error() {
let mut transport = StdioTransport::new(router());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut stdin_writer = server_stdin_writer;
stdin_writer.write_all(b"this is not json\n").await.unwrap();
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":42,\"method\":\"ping\"}\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let reader = BufReader::new(server_stdout_reader);
let frames = read_n_frames(reader, 2).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(
frames.len(),
2,
"expected parse-error frame + ping response, got: {frames:?}"
);
assert_jsonrpc_error_response(&frames[0]);
assert!(frames[0]["id"].is_null());
assert_eq!(frames[0]["error"]["code"].as_i64().unwrap(), -32700);
assert_eq!(frames[1]["jsonrpc"], "2.0");
assert_eq!(frames[1]["id"], 42);
assert!(
frames[1].get("result").is_some(),
"ping must return a successful result frame, got: {}",
frames[1]
);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn stdio_transport_preserves_partial_frame_when_read_is_cancelled() {
let mut transport = StdioTransport::new(router());
let control = transport.handle();
let (mut server_stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
server_stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":42,")
.await
.unwrap();
server_stdin_writer.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
control
.close_subscription(tower_mcp::protocol::RequestId::Number(999))
.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
server_stdin_writer
.write_all(b"\"method\":\"ping\"}\n")
.await
.unwrap();
drop(server_stdin_writer);
let frames = read_n_frames(BufReader::new(server_stdout_reader), 1).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(frames.len(), 1, "expected one response: {frames:?}");
assert_eq!(frames[0]["id"], 42, "the request prefix was lost");
assert!(
frames[0].get("result").is_some(),
"unexpected frame: {frames:?}"
);
}
#[tokio::test]
async fn stdio_transport_eof_returns_ok() {
let mut transport = StdioTransport::new(router());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, _server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
drop(server_stdin_writer);
let result = handle.await.expect("transport task join");
assert!(
result.is_ok(),
"run_with_streams must return Ok on EOF, got: {result:?}"
);
}
#[cfg(feature = "stateless")]
fn modern_router() -> McpRouter {
let inspect = ToolBuilder::new("inspect_meta")
.description("Report final request metadata visible to the handler")
.extractor_handler((), |ctx: Context, RawArgs(_): RawArgs| async move {
let version = ctx
.per_request_meta()
.and_then(|meta| meta.protocol_version.as_deref())
.unwrap_or("absent");
Ok(CallToolResult::text(format!(
"{version}|can_elicit={}",
ctx.can_elicit()
)))
})
.build();
McpRouter::new()
.server_info("stdio-modern-test", "0.0.0")
.tool(inspect)
}
#[cfg(feature = "stateless")]
fn final_meta(version: &str) -> serde_json::Value {
serde_json::json!({
"io.modelcontextprotocol/protocolVersion": version,
"io.modelcontextprotocol/clientInfo": {
"name": "stdio-loop-client",
"version": "1.0.0"
},
"io.modelcontextprotocol/clientCapabilities": {}
})
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn stdio_supports_final_and_legacy_lifecycles_on_one_stream() {
let mut transport = StdioTransport::new(modern_router());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(32 * 1024);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(32 * 1024);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let final_version = tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28;
let requests = vec![
serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": "server/discover",
"params": {"_meta": final_meta(final_version)}
}),
serde_json::json!({
"jsonrpc": "2.0", "id": 2, "method": "tools/list",
"params": {"_meta": final_meta(final_version)}
}),
serde_json::json!({
"jsonrpc": "2.0", "id": 3, "method": "tools/call",
"params": {
"name": "inspect_meta",
"arguments": {},
"_meta": final_meta(final_version)
}
}),
serde_json::json!({
"jsonrpc": "2.0", "id": 4, "method": "ping",
"params": {"_meta": final_meta(final_version)}
}),
serde_json::json!({
"jsonrpc": "2.0", "id": 5, "method": "server/discover",
"params": {"_meta": final_meta("2099-01-01")}
}),
serde_json::json!({
"jsonrpc": "2.0", "id": 6, "method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {"name": "legacy-client", "version": "1.0.0"}
}
}),
serde_json::json!({
"jsonrpc": "2.0", "method": "notifications/initialized"
}),
serde_json::json!({
"jsonrpc": "2.0", "id": 7, "method": "tools/list", "params": {}
}),
];
let mut stdin_writer = server_stdin_writer;
for request in requests {
stdin_writer
.write_all(format!("{request}\n").as_bytes())
.await
.unwrap();
}
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let frames = read_n_frames(BufReader::new(server_stdout_reader), 7).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(frames.len(), 7, "unexpected frames: {frames:#?}");
let by_id: std::collections::HashMap<i64, &serde_json::Value> = frames
.iter()
.filter_map(|f| f["id"].as_i64().map(|id| (id, f)))
.collect();
let response = |id: i64| {
*by_id
.get(&id)
.unwrap_or_else(|| panic!("no response for id {id} in {frames:#?}"))
};
let version_error = frames
.iter()
.find(|f| f["error"]["code"] == -32022)
.unwrap_or_else(|| panic!("no version-rejection error in {frames:#?}"));
assert_eq!(response(1)["result"]["resultType"], "complete");
assert_eq!(response(1)["result"]["ttlMs"], 0);
assert_eq!(response(1)["result"]["cacheScope"], "private");
assert_eq!(response(2)["result"]["resultType"], "complete");
assert_eq!(response(2)["result"]["ttlMs"], 0);
assert_eq!(
response(3)["result"]["content"][0]["text"],
format!("{final_version}|can_elicit=false")
);
assert_eq!(response(3)["result"]["resultType"], "complete");
assert!(response(3)["result"].get("ttlMs").is_none());
assert_eq!(response(4)["error"]["code"], -32601);
assert_eq!(version_error["error"]["data"]["requested"], "2099-01-01");
assert!(response(6)["result"].get("resultType").is_none());
assert!(response(7)["result"].get("resultType").is_none());
assert_eq!(response(7)["result"]["tools"].as_array().unwrap().len(), 1);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn stdio_runtime_protocol_allow_list_is_exact() {
let mut transport =
StdioTransport::new(modern_router()).protocol_support(ProtocolSupport::stable());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let request = serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": "server/discover",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28)
}
});
let mut stdin_writer = server_stdin_writer;
stdin_writer
.write_all(format!("{request}\n").as_bytes())
.await
.unwrap();
drop(stdin_writer);
let frames = read_n_frames(BufReader::new(server_stdout_reader), 1).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(frames[0]["error"]["code"], -32022);
assert_eq!(
frames[0]["error"]["data"]["requested"],
tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28
);
assert!(
frames[0]["error"]["data"]["supported"]
.as_array()
.unwrap()
.iter()
.all(|version| version != tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28)
);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn bidi_layer_preserves_runtime_protocol_allow_list() {
let mut transport = BidirectionalStdioTransport::new(modern_router())
.protocol_support(ProtocolSupport::stable())
.layer(tower::layer::util::Identity::new());
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let request = serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": "server/discover",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28)
}
});
stdin_writer
.write_all(format!("{request}\n").as_bytes())
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let mut reader = BufReader::new(server_stdout_reader);
let response = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("layered protocol rejection");
assert_eq!(response["error"]["code"], -32022);
assert_eq!(
response["error"]["data"]["requested"],
tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28
);
drop(stdin_writer);
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
}
#[cfg(feature = "stateless")]
fn subscription_router() -> McpRouter {
let emit = ToolBuilder::new("emit_changes")
.description("Emit every core subscription-scoped notification")
.extractor_handler((), |ctx: Context, RawArgs(_): RawArgs| async move {
ctx.notify_tools_list_changed();
ctx.notify_prompts_list_changed();
ctx.notify_resources_list_changed();
ctx.notify_resource_updated("file:///watched");
Ok(CallToolResult::text("emitted"))
})
.build();
McpRouter::new()
.server_info("stdio-subscription-test", "0.0.0")
.tool(emit)
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn stdio_multiplexes_and_gracefully_closes_final_subscriptions() {
let mut transport = StdioTransport::new(subscription_router());
let control = transport.handle();
let (server_stdin_writer, server_stdin) = tokio::io::duplex(32 * 1024);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(32 * 1024);
let transport_task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut writer = server_stdin_writer;
let mut reader = BufReader::new(server_stdout_reader);
let version = tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28;
let tools_listen = serde_json::json!({
"jsonrpc": "2.0",
"id": "tools-sub",
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(version),
"notifications": {
"toolsListChanged": true,
"promptsListChanged": false
}
}
});
writer
.write_all(format!("{tools_listen}\n").as_bytes())
.await
.unwrap();
writer.flush().await.unwrap();
let tools_ack = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("tools subscription acknowledgment");
assert_eq!(
tools_ack["method"],
"notifications/subscriptions/acknowledged"
);
assert_eq!(
tools_ack["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
"tools-sub"
);
assert_eq!(
tools_ack["params"]["notifications"]["toolsListChanged"],
true
);
assert!(
tools_ack["params"]["notifications"]
.get("promptsListChanged")
.is_none(),
"false filters must be omitted from the honored filter: {tools_ack}"
);
let resources_listen = serde_json::json!({
"jsonrpc": "2.0",
"id": 22,
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(version),
"notifications": {
"promptsListChanged": true,
"resourcesListChanged": true,
"resourceSubscriptions": ["file:///watched"]
}
}
});
writer
.write_all(format!("{resources_listen}\n").as_bytes())
.await
.unwrap();
writer.flush().await.unwrap();
let resources_ack = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("resources subscription acknowledgment");
assert_eq!(
resources_ack["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
22
);
let emit = serde_json::json!({
"jsonrpc": "2.0",
"id": 3,
"method": "tools/call",
"params": {
"name": "emit_changes",
"arguments": {},
"_meta": final_meta(version)
}
});
writer
.write_all(format!("{emit}\n").as_bytes())
.await
.unwrap();
writer.flush().await.unwrap();
let mut first_delivery = Vec::new();
for _ in 0..5 {
first_delivery.push(
timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("first subscription delivery"),
);
}
assert!(first_delivery.iter().any(|frame| frame["id"] == 3));
let delivered: Vec<_> = first_delivery
.iter()
.filter(|frame| frame.get("method").is_some())
.collect();
assert_eq!(
delivered.len(),
4,
"unexpected delivery: {first_delivery:#?}"
);
for frame in delivered {
let method = frame["method"].as_str().unwrap();
let id = &frame["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"];
match method {
"notifications/tools/list_changed" => assert_eq!(id, "tools-sub"),
"notifications/prompts/list_changed"
| "notifications/resources/list_changed"
| "notifications/resources/updated" => assert_eq!(id, 22),
_ => panic!("unexpected subscription notification: {frame}"),
}
}
let cancel = serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": {"requestId": "tools-sub"}
});
let emit_again = serde_json::json!({
"jsonrpc": "2.0",
"id": 4,
"method": "tools/call",
"params": {
"name": "emit_changes",
"arguments": {},
"_meta": final_meta(version)
}
});
writer
.write_all(format!("{cancel}\n{emit_again}\n").as_bytes())
.await
.unwrap();
writer.flush().await.unwrap();
let mut after_cancel = Vec::new();
for _ in 0..4 {
after_cancel.push(
timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("delivery after cancellation"),
);
}
assert!(after_cancel.iter().any(|frame| frame["id"] == 4));
assert!(
after_cancel
.iter()
.filter(|frame| frame.get("method").is_some())
.all(|frame| {
frame["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"]
== serde_json::json!(22)
})
);
control
.close_subscription(tower_mcp::RequestId::Number(22))
.unwrap();
let complete = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("graceful subscription result");
assert_eq!(complete["id"], 22);
assert_eq!(complete["result"]["resultType"], "complete");
assert_eq!(
complete["result"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
22
);
assert_eq!(
complete["result"]["_meta"]["io.modelcontextprotocol/serverInfo"]["name"],
"stdio-subscription-test"
);
let list = serde_json::json!({
"jsonrpc": "2.0",
"id": 5,
"method": "tools/list",
"params": {"_meta": final_meta(version)}
});
writer
.write_all(format!("{list}\n").as_bytes())
.await
.unwrap();
writer.flush().await.unwrap();
let list_response = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("shared channel remains usable");
assert_eq!(list_response["id"], 5);
let shutdown_listen = serde_json::json!({
"jsonrpc": "2.0",
"id": "shutdown-sub",
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(version),
"notifications": {"toolsListChanged": true}
}
});
writer
.write_all(format!("{shutdown_listen}\n").as_bytes())
.await
.unwrap();
writer.flush().await.unwrap();
let shutdown_ack = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("shutdown subscription acknowledgment");
assert_eq!(
shutdown_ack["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
"shutdown-sub"
);
control.shutdown().unwrap();
let shutdown_result = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("shutdown drains subscription result");
assert_eq!(shutdown_result["id"], "shutdown-sub");
assert_eq!(shutdown_result["result"]["resultType"], "complete");
transport_task
.await
.expect("transport task join")
.expect("run_with_streams ok");
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn stdio_subscription_validation_uses_runtime_protocol_policy() {
let mut transport =
StdioTransport::new(subscription_router()).protocol_support(ProtocolSupport::stable());
let (mut server_stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let request = serde_json::json!({
"jsonrpc": "2.0",
"id": 9,
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28),
"notifications": {"toolsListChanged": true}
}
});
server_stdin_writer
.write_all(format!("{request}\n").as_bytes())
.await
.unwrap();
drop(server_stdin_writer);
let frames = read_n_frames(BufReader::new(server_stdout_reader), 1).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(frames[0]["id"], 9);
assert_eq!(frames[0]["error"]["code"], -32022);
assert_eq!(
frames[0]["error"]["data"]["requested"],
tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28
);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn stdio_subscription_requires_final_metadata_and_filter() {
let mut transport = StdioTransport::new(subscription_router());
let (mut server_stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let version = tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28;
let missing_filter = serde_json::json!({
"jsonrpc": "2.0",
"id": 10,
"method": "subscriptions/listen",
"params": {"_meta": final_meta(version)}
});
let missing_capabilities = serde_json::json!({
"jsonrpc": "2.0",
"id": 11,
"method": "subscriptions/listen",
"params": {
"_meta": {
"io.modelcontextprotocol/protocolVersion": version
},
"notifications": {"toolsListChanged": true}
}
});
let legacy_shape = serde_json::json!({
"jsonrpc": "2.0",
"id": 12,
"method": "subscriptions/listen",
"params": {"notifications": {"toolsListChanged": true}}
});
server_stdin_writer
.write_all(format!("{missing_filter}\n{missing_capabilities}\n{legacy_shape}\n").as_bytes())
.await
.unwrap();
drop(server_stdin_writer);
let frames = read_n_frames(BufReader::new(server_stdout_reader), 3).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(frames[0]["id"], 10);
assert_eq!(frames[0]["error"]["code"], -32602);
assert_eq!(frames[1]["id"], 11);
assert_eq!(frames[1]["error"]["code"], -32602);
assert_eq!(frames[2]["id"], 12);
assert_eq!(
frames[2]["error"]["code"], -32600,
"claimless listen must remain on the legacy lifecycle and fail before initialize"
);
assert!(
frames.iter().all(|frame| frame.get("method").is_none()),
"invalid listens must not be acknowledged: {frames:#?}"
);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn middleware_wrapped_stdio_preserves_subscription_routing() {
let mut transport =
StdioTransport::new(subscription_router()).layer(tower::layer::util::Identity::new());
let control = transport.handle();
let (mut server_stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let request = serde_json::json!({
"jsonrpc": "2.0",
"id": "layered",
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28),
"notifications": {"toolsListChanged": true}
}
});
server_stdin_writer
.write_all(format!("{request}\n").as_bytes())
.await
.unwrap();
server_stdin_writer.flush().await.unwrap();
let mut reader = BufReader::new(server_stdout_reader);
let ack = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("layered subscription acknowledgment");
assert_eq!(
ack["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
"layered"
);
control
.close_subscription(tower_mcp::RequestId::String("layered".to_string()))
.unwrap();
let complete = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("layered graceful result");
assert_eq!(complete["id"], "layered");
assert_eq!(complete["result"]["resultType"], "complete");
assert!(
complete["result"]["_meta"]
.get("io.modelcontextprotocol/serverInfo")
.is_none(),
"generic services do not expose server identity to the transport"
);
control.shutdown().unwrap();
task.await
.expect("transport task join")
.expect("run_with_streams ok");
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn bidi_stdio_preserves_subscription_routing() {
let mut transport = BidirectionalStdioTransport::new(subscription_router())
.layer(tower::layer::util::Identity::new());
let control = transport.handle();
let (mut server_stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let request = serde_json::json!({
"jsonrpc": "2.0",
"id": 71,
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28),
"notifications": {"toolsListChanged": true}
}
});
server_stdin_writer
.write_all(format!("{request}\n").as_bytes())
.await
.unwrap();
server_stdin_writer.flush().await.unwrap();
let mut reader = BufReader::new(server_stdout_reader);
let ack = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("bidirectional subscription acknowledgment");
assert_eq!(
ack["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
71
);
let emit = serde_json::json!({
"jsonrpc": "2.0",
"id": 72,
"method": "tools/call",
"params": {
"name": "emit_changes",
"arguments": {},
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28)
}
});
server_stdin_writer
.write_all(format!("{emit}\n").as_bytes())
.await
.unwrap();
server_stdin_writer.flush().await.unwrap();
let deliveries = [
timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("bidirectional subscription delivery"),
timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("bidirectional tool response"),
];
let notification = deliveries
.iter()
.find(|frame| frame.get("method").is_some())
.expect("tagged notification must share the bidirectional stream");
assert_eq!(notification["method"], "notifications/tools/list_changed");
assert_eq!(
notification["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
71
);
assert!(deliveries.iter().any(|frame| frame["id"] == 72));
control
.close_subscription(tower_mcp::RequestId::Number(71))
.unwrap();
let complete = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("bidirectional graceful result");
assert_eq!(complete["id"], 71);
assert_eq!(complete["result"]["resultType"], "complete");
assert_eq!(
complete["result"]["_meta"]["io.modelcontextprotocol/serverInfo"]["name"],
"stdio-subscription-test"
);
control.shutdown().unwrap();
task.await
.expect("transport task join")
.expect("run_with_streams ok");
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn bidi_final_handlers_cannot_initiate_client_requests() {
let mut transport = BidirectionalStdioTransport::new(modern_router())
.layer(tower::layer::util::Identity::new());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let request = serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": "tools/call",
"params": {
"name": "inspect_meta",
"arguments": {},
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28)
}
});
let mut stdin_writer = server_stdin_writer;
stdin_writer
.write_all(format!("{request}\n").as_bytes())
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let mut reader = BufReader::new(server_stdout_reader);
let frame = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("final tool response");
assert_eq!(
frame["result"]["content"][0]["text"],
"2026-07-28|can_elicit=false"
);
assert_eq!(frame["result"]["resultType"], "complete");
drop(stdin_writer);
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
}
#[tokio::test]
async fn bidi_transport_parse_error_wire_shape() {
let mut transport = BidirectionalStdioTransport::new(router());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut stdin_writer = server_stdin_writer;
stdin_writer
.write_all(b"not valid json{{{\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let reader = BufReader::new(server_stdout_reader);
let frames = read_n_frames(reader, 1).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(
frames.len(),
1,
"expected one parse-error frame, got: {frames:?}"
);
assert_jsonrpc_error_response(&frames[0]);
assert!(frames[0]["id"].is_null());
assert_eq!(frames[0]["error"]["code"].as_i64().unwrap(), -32700);
}
#[tokio::test]
async fn bidi_transport_loop_continues_after_parse_error() {
let mut transport = BidirectionalStdioTransport::new(router());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut stdin_writer = server_stdin_writer;
stdin_writer.write_all(b"this is not json\n").await.unwrap();
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":7,\"method\":\"ping\"}\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let reader = BufReader::new(server_stdout_reader);
let frames = read_n_frames(reader, 2).await;
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(
frames.len(),
2,
"expected parse-error frame + ping response, got: {frames:?}"
);
assert_jsonrpc_error_response(&frames[0]);
assert!(frames[0]["id"].is_null());
assert_eq!(frames[0]["error"]["code"].as_i64().unwrap(), -32700);
assert_eq!(frames[1]["jsonrpc"], "2.0");
assert_eq!(frames[1]["id"], 7);
assert!(
frames[1].get("result").is_some(),
"ping must return a successful result frame, got: {}",
frames[1]
);
}
#[derive(Debug, Deserialize, JsonSchema)]
struct NoArgs {}
async fn read_frame<R>(reader: &mut BufReader<R>) -> serde_json::Value
where
R: tokio::io::AsyncRead + Unpin,
{
loop {
let mut line = String::new();
let n = reader
.read_line(&mut line)
.await
.expect("read from server output");
if n == 0 {
panic!("EOF before a frame was read");
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
return serde_json::from_str(trimmed)
.unwrap_or_else(|e| panic!("invalid JSON on output: {e}: {trimmed}"));
}
}
#[tokio::test]
async fn bidi_transport_wires_client_requester_into_context() {
let check = ToolBuilder::new("check_elicit")
.description("Report whether elicitation is available")
.extractor_handler((), |ctx: Context, Json(_): Json<NoArgs>| async move {
Ok(CallToolResult::text(ctx.can_elicit().to_string()))
})
.build();
let router = McpRouter::new()
.server_info("bidi-elicit-test", "0.0.0")
.tool(check);
let mut transport = BidirectionalStdioTransport::new(router);
let (server_stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut w = server_stdin_writer;
w.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"initialize\",\"params\":{\"protocolVersion\":\"2025-11-25\",\"capabilities\":{\"elicitation\":{}},\"clientInfo\":{\"name\":\"t\",\"version\":\"0\"}}}\n").await.unwrap();
w.write_all(b"{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\"}\n")
.await
.unwrap();
w.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":2,\"method\":\"tools/call\",\"params\":{\"name\":\"check_elicit\",\"arguments\":{}}}\n").await.unwrap();
w.flush().await.unwrap();
let reader = BufReader::new(server_stdout_reader);
let frames = read_n_frames(reader, 2).await;
drop(w);
let _ = handle.await;
let call = &frames[1];
assert_eq!(call["id"], 2, "expected tools/call response, got: {call}");
let text = call["result"]["content"][0]["text"].as_str().unwrap_or("");
assert_eq!(
text, "true",
"can_elicit() must be true once the requester is wired, got: {call}"
);
}
#[tokio::test]
async fn bidi_transport_elicitation_round_trip() {
let confirm = ToolBuilder::new("confirm")
.description("Confirm an action via elicitation")
.extractor_handler((), |ctx: Context, Json(_): Json<NoArgs>| async move {
let params = ElicitFormParams {
message: "Confirm?".to_string(),
requested_schema: ElicitFormSchema::new().boolean_field(
"confirmed",
Some("Confirm"),
true,
),
mode: None,
meta: None,
};
match ctx.elicit_form(params).await {
Ok(result) => {
let accepted = matches!(result.action, ElicitAction::Accept);
Ok(CallToolResult::text(format!("confirmed={accepted}")))
}
Err(e) => Ok(CallToolResult::error(format!("elicit failed: {e}"))),
}
})
.build();
let router = McpRouter::new()
.server_info("bidi-elicit-test", "0.0.0")
.tool(confirm);
let mut transport =
BidirectionalStdioTransport::new(router).layer(tower::layer::util::Identity::new());
let (server_stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let driver = async move {
let mut w = server_stdin_writer;
let mut reader = BufReader::new(server_stdout_reader);
w.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"initialize\",\"params\":{\"protocolVersion\":\"2025-11-25\",\"capabilities\":{\"elicitation\":{}},\"clientInfo\":{\"name\":\"t\",\"version\":\"0\"}}}\n").await.unwrap();
w.flush().await.unwrap();
let init = read_frame(&mut reader).await;
assert_eq!(init["id"], 1, "expected initialize response, got: {init}");
w.write_all(b"{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\"}\n")
.await
.unwrap();
w.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":2,\"method\":\"tools/call\",\"params\":{\"name\":\"confirm\",\"arguments\":{}}}\n").await.unwrap();
w.flush().await.unwrap();
loop {
let frame = read_frame(&mut reader).await;
if frame.get("method").and_then(|m| m.as_str()) == Some("elicitation/create") {
let response = serde_json::json!({
"jsonrpc": "2.0",
"id": frame["id"],
"result": { "action": "accept", "content": { "confirmed": true } }
});
w.write_all(format!("{response}\n").as_bytes())
.await
.unwrap();
w.flush().await.unwrap();
continue;
}
if frame["id"] == serde_json::json!(2) {
assert!(
frame.get("result").is_some(),
"tools/call must succeed after elicitation, got: {frame}"
);
let text = frame["result"]["content"][0]["text"].as_str().unwrap_or("");
assert_eq!(
text, "confirmed=true",
"elicitation round-trip should return the client's accept, got: {frame}"
);
break;
}
}
drop(w);
};
timeout(Duration::from_secs(5), driver)
.await
.expect("elicitation round-trip over bidirectional stdio timed out (deadlock)");
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn bidi_layer_observes_an_inbound_tool_request() {
use std::sync::{Arc, Mutex};
use tower::ServiceExt as _;
use tower_mcp::router::RouterRequest;
let observed = Arc::new(Mutex::new(Vec::new()));
let layer = tower::layer::layer_fn({
let observed = observed.clone();
move |inner: McpRouter| {
let observed = observed.clone();
tower::service_fn(move |request: RouterRequest| {
let inner = inner.clone();
let observed = observed.clone();
async move {
observed
.lock()
.expect("observation lock")
.push(request.inner.method_name().to_string());
inner.oneshot(request).await
}
})
}
});
let tool = ToolBuilder::new("observed")
.description("Return a fixed result")
.handler(|_: serde_json::Value| async { Ok(CallToolResult::text("seen")) })
.build();
let router = McpRouter::new()
.server_info("bidi-layer-test", "0.0.0")
.tool(tool);
let mut transport = BidirectionalStdioTransport::new(router).layer(layer);
let (mut stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut reader = BufReader::new(server_stdout_reader);
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"initialize\",\"params\":{\"protocolVersion\":\"2025-11-25\",\"capabilities\":{},\"clientInfo\":{\"name\":\"t\",\"version\":\"0\"}}}\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let init = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("layered initialize response");
assert_eq!(init["id"], 1);
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\"}\n")
.await
.unwrap();
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":2,\"method\":\"tools/call\",\"params\":{\"name\":\"observed\",\"arguments\":{}}}\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let call = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("layered tool response");
assert_eq!(call["id"], 2);
assert_eq!(call["result"]["content"][0]["text"], "seen");
drop(stdin_writer);
task.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(
*observed.lock().expect("observation lock"),
["initialize", "tools/call"]
);
}
#[derive(Clone)]
struct RejectFirstReadiness {
inner: McpRouter,
reject_next: std::sync::Arc<std::sync::atomic::AtomicBool>,
}
impl tower::Service<tower_mcp::router::RouterRequest> for RejectFirstReadiness {
type Response = tower_mcp::router::RouterResponse;
type Error = std::io::Error;
type Future = std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Self::Response, Self::Error>> + Send>,
>;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
use std::sync::atomic::Ordering;
if self.reject_next.swap(false, Ordering::SeqCst) {
return std::task::Poll::Ready(Err(std::io::Error::other("middleware was not ready")));
}
tower::Service::poll_ready(&mut self.inner, cx).map_err(|never| match never {})
}
fn call(&mut self, request: tower_mcp::router::RouterRequest) -> Self::Future {
let future = tower::Service::call(&mut self.inner, request);
Box::pin(async move { Ok(future.await.expect("MCP router service is infallible")) })
}
}
#[tokio::test]
async fn bidi_layer_readiness_errors_become_framed_jsonrpc_errors() {
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
let reject_next = Arc::new(AtomicBool::new(true));
let layer = tower::layer::layer_fn(move |inner: McpRouter| RejectFirstReadiness {
inner,
reject_next: reject_next.clone(),
});
let mut transport = BidirectionalStdioTransport::new(router()).layer(layer);
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut reader = BufReader::new(server_stdout_reader);
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":7,\"method\":\"ping\"}\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let response = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("layer readiness error response");
assert_jsonrpc_error_response(&response);
assert_eq!(response["id"], 7);
assert_eq!(response["error"]["code"], -32603);
assert_eq!(response["error"]["message"], "middleware was not ready");
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":8,\"method\":\"ping\"}\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let response = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("request after layer readiness error response");
assert_eq!(response["id"], 8);
assert!(response.get("result").is_some());
drop(stdin_writer);
task.await
.expect("transport task join")
.expect("run_with_streams ok");
}
#[tokio::test]
async fn bidi_layer_errors_become_framed_jsonrpc_errors() {
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tower::ServiceExt as _;
use tower_mcp::router::{RouterRequest, RouterResponse};
let reject_next = Arc::new(AtomicBool::new(true));
let layer = tower::layer::layer_fn(move |inner: McpRouter| {
let reject_next = reject_next.clone();
tower::service_fn(move |request: RouterRequest| {
let inner = inner.clone();
let reject = reject_next.swap(false, Ordering::SeqCst);
async move {
if reject {
return Err::<RouterResponse, _>(std::io::Error::other(
"middleware rejected request",
));
}
Ok(inner
.oneshot(request)
.await
.expect("MCP router service is infallible"))
}
})
});
let mut transport = BidirectionalStdioTransport::new(router()).layer(layer);
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let mut reader = BufReader::new(server_stdout_reader);
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":9,\"method\":\"ping\"}\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let response = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("layer error response");
assert_jsonrpc_error_response(&response);
assert_eq!(response["id"], 9);
assert_eq!(response["error"]["code"], -32603);
assert_eq!(response["error"]["message"], "middleware rejected request");
stdin_writer
.write_all(b"{\"jsonrpc\":\"2.0\",\"id\":10,\"method\":\"ping\"}\n")
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let response = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("request after layer error response");
assert_eq!(response["id"], 10);
assert!(response.get("result").is_some());
drop(stdin_writer);
task.await
.expect("transport task join")
.expect("run_with_streams ok");
}
#[cfg(feature = "stateless")]
mod listen_observation {
use super::*;
use std::sync::{Arc, Mutex};
use std::task::{Context as TaskContext, Poll};
use tower_mcp::router::{RouterRequest, RouterResponse};
#[derive(Clone)]
struct RecordingLayer {
seen: Arc<Mutex<Vec<String>>>,
}
#[derive(Clone)]
struct RecordingService<S> {
inner: S,
seen: Arc<Mutex<Vec<String>>>,
}
impl<S> tower::Layer<S> for RecordingLayer {
type Service = RecordingService<S>;
fn layer(&self, inner: S) -> Self::Service {
RecordingService {
inner,
seen: self.seen.clone(),
}
}
}
impl<S> tower_service::Service<RouterRequest> for RecordingService<S>
where
S: tower_service::Service<RouterRequest, Response = RouterResponse>,
{
type Response = S::Response;
type Error = S::Error;
type Future = S::Future;
fn poll_ready(&mut self, cx: &mut TaskContext<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: RouterRequest) -> Self::Future {
self.seen
.lock()
.unwrap()
.push(request.inner.method_name().to_string());
self.inner.call(request)
}
}
async fn run_listen_exchange(request: serde_json::Value) -> (Vec<String>, serde_json::Value) {
let seen = Arc::new(Mutex::new(Vec::new()));
let mut transport =
StdioTransport::new(subscription_router()).layer(RecordingLayer { seen: seen.clone() });
let control = transport.handle();
let (mut server_stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(8192);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
server_stdin_writer
.write_all(format!("{request}\n").as_bytes())
.await
.unwrap();
server_stdin_writer.flush().await.unwrap();
let mut reader = BufReader::new(server_stdout_reader);
let frame = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("listen exchange reply");
control.shutdown().unwrap();
let _ = task.await;
let observed = seen.lock().unwrap().clone();
(observed, frame)
}
#[tokio::test]
async fn middleware_observes_an_accepted_listen() {
let (observed, frame) = run_listen_exchange(serde_json::json!({
"jsonrpc": "2.0",
"id": "observed",
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28),
"notifications": {"toolsListChanged": true}
}
}))
.await;
assert_eq!(observed, vec!["subscriptions/listen"]);
assert_eq!(frame["method"], "notifications/subscriptions/acknowledged");
assert_eq!(
frame["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
"observed"
);
assert_eq!(frame["params"]["notifications"]["toolsListChanged"], true);
}
#[tokio::test]
async fn middleware_observes_a_rejected_listen() {
let (observed, frame) = run_listen_exchange(serde_json::json!({
"jsonrpc": "2.0",
"id": "rejected",
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28),
"notifications": {"taskIds": ["task-1"]}
}
}))
.await;
assert_eq!(observed, vec!["subscriptions/listen"]);
assert_eq!(frame["id"], "rejected");
assert_eq!(frame["error"]["code"], -32021);
assert!(
frame["error"]["data"]["requiredCapabilities"]["extensions"]
["io.modelcontextprotocol/tasks"]
.is_object(),
"the rejection must name the missing extension: {frame}"
);
}
#[tokio::test]
async fn schema_rejections_stay_ahead_of_the_middleware_boundary() {
let (observed, frame) = run_listen_exchange(serde_json::json!({
"jsonrpc": "2.0",
"id": "no-filter",
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28)
}
}))
.await;
assert!(observed.is_empty(), "schema rejections precede middleware");
assert_eq!(frame["id"], "no-filter");
assert_eq!(frame["error"]["code"], -32602);
assert!(
frame["error"]["message"]
.as_str()
.unwrap_or_default()
.contains("required `notifications` field is missing"),
"the rejection must explain the missing filter: {frame}"
);
}
}
#[cfg(feature = "stateless")]
mod close_observation {
use super::*;
use std::sync::{Arc, Mutex};
use tower_mcp::{SubscriptionClose, SubscriptionCloseReason, SubscriptionObserver};
#[derive(Default)]
struct RecordingObserver {
closes: Mutex<Vec<SubscriptionClose>>,
}
impl SubscriptionObserver for RecordingObserver {
fn on_close(&self, close: SubscriptionClose) {
self.closes.lock().unwrap().push(close);
}
}
fn observed_router(observer: Arc<RecordingObserver>) -> McpRouter {
subscription_router().with_subscription_observer(observer)
}
fn listen_frame(id: &str) -> serde_json::Value {
serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"method": "subscriptions/listen",
"params": {
"_meta": final_meta(tower_mcp::protocol::PROTOCOL_VERSION_2026_07_28),
"notifications": {"toolsListChanged": true}
}
})
}
#[tokio::test]
async fn cancellation_reaches_the_observer() {
let observer = Arc::new(RecordingObserver::default());
let mut transport = StdioTransport::new(observed_router(observer.clone()));
let control = transport.handle();
let (mut stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, stdout_reader) = tokio::io::duplex(8192);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
stdin_writer
.write_all(format!("{}\n", listen_frame("cancel-me")).as_bytes())
.await
.unwrap();
let mut reader = BufReader::new(stdout_reader);
let _ack = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("acknowledgment");
let cancel = serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": {"requestId": "cancel-me"}
});
stdin_writer
.write_all(format!("{cancel}\n").as_bytes())
.await
.unwrap();
stdin_writer.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
control.shutdown().unwrap();
let _ = task.await;
let closes = observer.closes.lock().unwrap();
assert_eq!(closes.len(), 1, "exactly one close record");
assert_eq!(
closes[0].subscription_id,
tower_mcp::RequestId::String("cancel-me".to_string())
);
assert_eq!(closes[0].reason, SubscriptionCloseReason::Cancelled);
}
#[tokio::test]
async fn graceful_close_reports_drained() {
let observer = Arc::new(RecordingObserver::default());
let mut transport = StdioTransport::new(observed_router(observer.clone()));
let control = transport.handle();
let (mut stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, stdout_reader) = tokio::io::duplex(8192);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
stdin_writer
.write_all(format!("{}\n", listen_frame("drain-me")).as_bytes())
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let mut reader = BufReader::new(stdout_reader);
let _ack = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("acknowledgment");
control
.close_subscription(tower_mcp::RequestId::String("drain-me".to_string()))
.unwrap();
let complete = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("graceful terminal result");
assert_eq!(complete["result"]["resultType"], "complete");
control.shutdown().unwrap();
let _ = task.await;
let closes = observer.closes.lock().unwrap();
assert_eq!(closes.len(), 1);
assert_eq!(closes[0].reason, SubscriptionCloseReason::Drained);
}
#[tokio::test]
async fn eof_reports_disconnected() {
let observer = Arc::new(RecordingObserver::default());
let mut transport = StdioTransport::new(observed_router(observer.clone()));
let (mut stdin_writer, server_stdin) = tokio::io::duplex(8192);
let (server_stdout, stdout_reader) = tokio::io::duplex(8192);
let task = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
stdin_writer
.write_all(format!("{}\n", listen_frame("drop-me")).as_bytes())
.await
.unwrap();
stdin_writer.flush().await.unwrap();
let mut reader = BufReader::new(stdout_reader);
let _ack = timeout(Duration::from_secs(2), read_frame(&mut reader))
.await
.expect("acknowledgment");
drop(stdin_writer);
timeout(Duration::from_secs(2), task)
.await
.expect("loop must end at EOF")
.expect("join")
.expect("run ok");
let closes = observer.closes.lock().unwrap();
assert_eq!(closes.len(), 1);
assert_eq!(
closes[0].subscription_id,
tower_mcp::RequestId::String("drop-me".to_string())
);
assert_eq!(closes[0].reason, SubscriptionCloseReason::Disconnected);
}
}
mod concurrency {
use super::*;
use tower_mcp::extract::RawArgs;
const SLOW: Duration = Duration::from_millis(400);
fn router_with_a_slow_tool() -> McpRouter {
let slow = ToolBuilder::new("slow")
.description("Sleeps before answering")
.extractor_handler((), |_ctx: Context, RawArgs(_): RawArgs| async move {
tokio::time::sleep(SLOW).await;
Ok(CallToolResult::text("slow"))
})
.build();
let fast = ToolBuilder::new("fast")
.description("Answers immediately")
.extractor_handler((), |_ctx: Context, RawArgs(_): RawArgs| async move {
Ok(CallToolResult::text("fast"))
})
.build();
McpRouter::new()
.server_info("stdio-concurrency-test", "0.0.0")
.tool(slow)
.tool(fast)
}
const INIT: &str = r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"t","version":"1"}}}"#;
const CALL_SLOW: &str =
r#"{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"slow","arguments":{}}}"#;
const CALL_FAST: &str =
r#"{"jsonrpc":"2.0","id":3,"method":"tools/call","params":{"name":"fast","arguments":{}}}"#;
async fn response_id_order(mut transport: StdioTransport) -> Vec<i64> {
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let handle = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
for line in [INIT, CALL_SLOW, CALL_FAST] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let frames = timeout(
Duration::from_secs(10),
read_n_frames(BufReader::new(server_stdout_reader), 3),
)
.await
.expect("responses must arrive");
handle
.await
.expect("transport task join")
.expect("run_with_streams ok");
assert_eq!(frames.len(), 3, "expected three responses, got {frames:?}");
frames
.iter()
.map(|f| f["id"].as_i64().expect("response carries an id"))
.collect()
}
#[tokio::test]
async fn a_slow_tool_does_not_block_later_requests() {
let ids = response_id_order(StdioTransport::new(router_with_a_slow_tool())).await;
assert_eq!(
ids,
vec![1, 3, 2],
"the fast call (3) was issued after the slow one (2) and must be answered first"
);
}
#[tokio::test]
async fn a_limit_of_one_restores_serial_handling() {
let transport = StdioTransport::new(router_with_a_slow_tool()).max_concurrent_requests(1);
let ids = response_id_order(transport).await;
assert_eq!(
ids,
vec![1, 2, 3],
"with a limit of 1 the slow call (2) must be answered before the fast one (3)"
);
}
}
mod control_under_saturation {
use super::*;
use tower_mcp::extract::RawArgs;
fn router_with_a_waiting_tool() -> McpRouter {
let wait = ToolBuilder::new("wait")
.description("Waits until cancelled")
.extractor_handler((), |ctx: Context, RawArgs(_): RawArgs| async move {
ctx.cancelled().await;
Ok(CallToolResult::text("cancelled"))
})
.build();
McpRouter::new()
.server_info("stdio-control-test", "0.0.0")
.tool(wait)
}
const INIT: &str = r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"t","version":"1"}}}"#;
const CALL_WAIT: &str =
r#"{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"wait","arguments":{}}}"#;
const PING: &str = r#"{"jsonrpc":"2.0","id":3,"method":"ping","params":{}}"#;
const CANCEL: &str =
r#"{"jsonrpc":"2.0","method":"notifications/cancelled","params":{"requestId":2}}"#;
#[tokio::test]
async fn cancellation_reaches_a_running_handler_without_a_limit() {
let mut transport = StdioTransport::new(router_with_a_waiting_tool());
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
for line in [INIT, CALL_WAIT] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(150)).await;
for line in [PING, CANCEL] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
let frames = timeout(
Duration::from_secs(5),
read_n_frames(BufReader::new(server_stdout_reader), 3),
)
.await
.expect("unlimited transport must answer all three");
let ids: Vec<i64> = frames.iter().filter_map(|f| f["id"].as_i64()).collect();
assert!(ids.contains(&2), "cancelled call answered: {frames:?}");
}
#[tokio::test]
async fn a_saturated_limit_must_not_starve_cancellation() {
let mut transport =
StdioTransport::new(router_with_a_waiting_tool()).max_concurrent_requests(1);
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
for line in [INIT, CALL_WAIT] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(150)).await;
for line in [PING, CANCEL] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
let frames = timeout(
Duration::from_secs(5),
read_n_frames(BufReader::new(server_stdout_reader), 3),
)
.await
.expect("cancellation must be readable while the limit is saturated");
let ids: Vec<i64> = frames.iter().filter_map(|f| f["id"].as_i64()).collect();
assert!(
ids.contains(&2),
"the cancelled call must be answered: {frames:?}"
);
assert!(
ids.contains(&3),
"the queued request must run once the permit frees: {frames:?}"
);
}
}
mod inbound_notifications {
use super::*;
use tower_mcp::extract::RawArgs;
fn router_with_a_waiting_tool() -> McpRouter {
let wait = ToolBuilder::new("wait")
.description("Waits until cancelled")
.extractor_handler((), |ctx: Context, RawArgs(_): RawArgs| async move {
ctx.cancelled().await;
Ok(CallToolResult::text("cancelled"))
})
.build();
McpRouter::new()
.server_info("inbound-test", "0.0.0")
.tool(wait)
}
const INIT: &str = r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"t","version":"1"}}}"#;
const CALL_WAIT: &str =
r#"{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"wait","arguments":{}}}"#;
const CANCEL: &str =
r#"{"jsonrpc":"2.0","method":"notifications/cancelled","params":{"requestId":2}}"#;
async fn cancel_a_running_call<F, Fut>(run: F) -> Vec<serde_json::Value>
where
F: FnOnce(tokio::io::DuplexStream, tokio::io::DuplexStream) -> Fut,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
tokio::spawn(run(server_stdin, server_stdout));
for line in [INIT, CALL_WAIT] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(150)).await;
stdin_writer.write_all(CANCEL.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
stdin_writer.flush().await.unwrap();
timeout(
Duration::from_secs(5),
read_n_frames(BufReader::new(server_stdout_reader), 2),
)
.await
.expect("the cancelled call must be answered")
}
#[tokio::test]
async fn cancellation_survives_a_middleware_layer() {
let frames = cancel_a_running_call(|stdin, stdout| async move {
let mut transport = StdioTransport::new(router_with_a_waiting_tool())
.layer(tower::layer::util::Identity::new());
let _ = transport.run_with_streams(stdin, stdout).await;
})
.await;
let answer = frames
.iter()
.find(|f| f["id"] == 2)
.unwrap_or_else(|| panic!("no answer for the cancelled call: {frames:?}"));
assert_eq!(answer["result"]["content"][0]["text"], "cancelled");
}
#[tokio::test]
async fn cancellation_works_without_server_notifications() {
let frames = cancel_a_running_call(|stdin, stdout| async move {
let mut transport =
StdioTransport::without_server_notifications(router_with_a_waiting_tool());
let _ = transport.run_with_streams(stdin, stdout).await;
})
.await;
let answer = frames
.iter()
.find(|f| f["id"] == 2)
.unwrap_or_else(|| panic!("no answer for the cancelled call: {frames:?}"));
assert_eq!(answer["result"]["content"][0]["text"], "cancelled");
}
#[tokio::test]
async fn declining_outbound_notifications_drops_the_logging_capability() {
async fn initialize_with(transport: impl FnOnce() -> StdioTransport) -> serde_json::Value {
let mut transport = transport();
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
stdin_writer.write_all(INIT.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
read_n_frames(BufReader::new(server_stdout_reader), 1)
.await
.remove(0)
}
let with = initialize_with(|| StdioTransport::new(router_with_a_waiting_tool())).await;
assert!(
with["result"]["capabilities"]["logging"].is_object(),
"the default still advertises logging: {with}"
);
assert_eq!(with["result"]["capabilities"]["tools"]["listChanged"], true);
let without = initialize_with(|| {
StdioTransport::without_server_notifications(router_with_a_waiting_tool())
})
.await;
assert!(
without["result"]["capabilities"]["logging"].is_null(),
"logging must not be advertised when nothing can be sent: {without}"
);
assert_ne!(
without["result"]["capabilities"]["tools"]["listChanged"], true,
"listChanged must not be claimed either: {without}"
);
}
}
mod shutdown {
use super::*;
use std::sync::Arc;
use tokio::sync::Notify;
use tower_mcp::extract::RawArgs;
fn router_needing_app_shutdown(app_stopped: Arc<Notify>) -> McpRouter {
let wait = ToolBuilder::new("wait")
.description("Finishes only on application shutdown")
.extractor_handler((), move |_ctx: Context, RawArgs(_): RawArgs| {
let app_stopped = app_stopped.clone();
async move {
app_stopped.notified().await;
Ok(CallToolResult::text("stopped"))
}
})
.build();
McpRouter::new()
.server_info("shutdown-test", "0.0.0")
.tool(wait)
}
const INIT: &str = r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"t","version":"1"}}}"#;
const CALL_WAIT: &str =
r#"{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"wait","arguments":{}}}"#;
#[tokio::test]
async fn stopping_fires_before_the_drain_so_the_application_can_release_handlers() {
let app_stopped = Arc::new(Notify::new());
let mut transport = StdioTransport::new(router_needing_app_shutdown(app_stopped.clone()));
let handle = transport.handle();
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
let server = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let release = app_stopped.clone();
tokio::spawn(async move {
handle.stopping().await;
release.notify_waiters();
});
for line in [INIT, CALL_WAIT] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(150)).await;
drop(stdin_writer);
let frames = timeout(
Duration::from_secs(5),
read_n_frames(BufReader::new(server_stdout_reader), 2),
)
.await
.expect("run must not deadlock against application shutdown");
timeout(Duration::from_secs(5), server)
.await
.expect("run must return")
.expect("join")
.expect("run_with_streams ok");
let answer = frames
.iter()
.find(|f| f["id"] == 2)
.unwrap_or_else(|| panic!("the released call must be answered: {frames:?}"));
assert_eq!(answer["result"]["content"][0]["text"], "stopped");
}
#[tokio::test]
async fn a_drain_deadline_lets_run_return_when_a_handler_never_finishes() {
let never = Arc::new(Notify::new());
let mut transport = StdioTransport::new(router_needing_app_shutdown(never))
.drain_timeout(Duration::from_millis(200));
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, _reader) = tokio::io::duplex(4096);
let server = tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
for line in [INIT, CALL_WAIT] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(150)).await;
drop(stdin_writer);
timeout(Duration::from_secs(5), server)
.await
.expect("the deadline must let run return")
.expect("join")
.expect("run_with_streams ok");
}
#[tokio::test]
async fn the_default_still_waits_for_a_finite_call() {
let slow = ToolBuilder::new("slow")
.description("Takes a moment")
.extractor_handler((), |_ctx: Context, RawArgs(_): RawArgs| async move {
tokio::time::sleep(Duration::from_millis(300)).await;
Ok(CallToolResult::text("done"))
})
.build();
let router = McpRouter::new()
.server_info("shutdown-test", "0.0.0")
.tool(slow);
let mut transport = StdioTransport::new(router);
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let call = r#"{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"slow","arguments":{}}}"#;
for line in [INIT, call] {
stdin_writer.write_all(line.as_bytes()).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
}
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let frames = timeout(
Duration::from_secs(5),
read_n_frames(BufReader::new(server_stdout_reader), 2),
)
.await
.expect("the in-flight call must still be answered");
let answer = frames.iter().find(|f| f["id"] == 2).expect("answered");
assert_eq!(answer["result"]["content"][0]["text"], "done");
}
#[tokio::test]
async fn an_explicit_shutdown_also_fires_stopping() {
let mut transport = StdioTransport::new(McpRouter::new().server_info("s", "0.0.0"));
let handle = transport.handle();
let (_stdin_writer, server_stdin) = tokio::io::duplex(64);
let (server_stdout, _reader) = tokio::io::duplex(64);
tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let observer = handle.clone();
let observed = tokio::spawn(async move { observer.stopping().await });
tokio::time::sleep(Duration::from_millis(50)).await;
handle.shutdown().expect("shutdown");
timeout(Duration::from_secs(5), observed)
.await
.expect("stopping must fire on an explicit shutdown")
.expect("join");
}
}
mod invalid_utf8 {
use super::*;
use tower_mcp::GenericStdioTransport;
const PING: &[u8] = br#"{"jsonrpc":"2.0","id":42,"method":"ping"}"#;
async fn answers_after_a_bad_byte<F, Fut>(run: F) -> Vec<serde_json::Value>
where
F: FnOnce(tokio::io::DuplexStream, tokio::io::DuplexStream) -> Fut,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
tokio::spawn(run(server_stdin, server_stdout));
stdin_writer.write_all(&[0xff, 0xfe, b'\n']).await.unwrap();
stdin_writer.write_all(PING).await.unwrap();
stdin_writer.write_all(b"\n").await.unwrap();
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
timeout(
Duration::from_secs(5),
read_n_frames(BufReader::new(server_stdout_reader), 2),
)
.await
.expect("the transport must survive the bad byte and answer the ping")
}
fn assert_parse_error_then_ping(frames: &[serde_json::Value]) {
assert_eq!(
frames.len(),
2,
"expected a parse error and a ping answer, got: {frames:?}"
);
assert_jsonrpc_error_response(&frames[0]);
assert!(
frames[0]["id"].is_null(),
"a discarded frame has no recoverable id: {}",
frames[0]
);
assert_eq!(frames[0]["error"]["code"].as_i64().unwrap(), -32700);
assert_eq!(
frames[1]["id"], 42,
"the request after the bad byte must be served: {}",
frames[1]
);
assert!(
frames[1].get("result").is_some(),
"ping must return a successful result frame, got: {}",
frames[1]
);
}
#[tokio::test]
async fn stdio_transport_answers_and_keeps_serving() {
let frames = answers_after_a_bad_byte(|stdin, stdout| async move {
let mut transport = StdioTransport::new(router());
let _ = transport.run_with_streams(stdin, stdout).await;
})
.await;
assert_parse_error_then_ping(&frames);
}
#[tokio::test]
async fn a_layered_transport_answers_and_keeps_serving() {
let frames = answers_after_a_bad_byte(|stdin, stdout| async move {
let mut transport =
StdioTransport::new(router()).layer(tower::layer::util::Identity::new());
let _ = transport.run_with_streams(stdin, stdout).await;
})
.await;
assert_parse_error_then_ping(&frames);
}
#[tokio::test]
async fn a_generic_transport_without_notifications_answers_and_keeps_serving() {
let frames = answers_after_a_bad_byte(|stdin, stdout| async move {
let mut transport = GenericStdioTransport::new(router());
let _ = transport.run_with_streams(stdin, stdout).await;
})
.await;
assert_parse_error_then_ping(&frames);
}
#[tokio::test]
async fn bidi_transport_answers_and_keeps_serving() {
let frames = answers_after_a_bad_byte(|stdin, stdout| async move {
let mut transport = BidirectionalStdioTransport::new(router());
let _ = transport.run_with_streams(stdin, stdout).await;
})
.await;
assert_parse_error_then_ping(&frames);
}
#[tokio::test]
async fn a_multibyte_character_split_across_reads_still_decodes() {
let mut transport = StdioTransport::new(router());
let (mut stdin_writer, server_stdin) = tokio::io::duplex(4096);
let (server_stdout, server_stdout_reader) = tokio::io::duplex(4096);
tokio::spawn(async move {
transport
.run_with_streams(server_stdin, server_stdout)
.await
});
let frame = "{\"jsonrpc\":\"2.0\",\"id\":\"caf\u{e9}\",\"method\":\"ping\"}\n";
let split = frame.find('\u{e9}').expect("the frame carries the char") + 1;
stdin_writer
.write_all(&frame.as_bytes()[..split])
.await
.unwrap();
stdin_writer.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
stdin_writer
.write_all(&frame.as_bytes()[split..])
.await
.unwrap();
stdin_writer.flush().await.unwrap();
drop(stdin_writer);
let frames = timeout(
Duration::from_secs(5),
read_n_frames(BufReader::new(server_stdout_reader), 1),
)
.await
.expect("the split frame must be answered");
assert_eq!(frames.len(), 1, "expected one answer, got: {frames:?}");
assert_eq!(
frames[0]["id"], "caf\u{e9}",
"the character must survive the split: {}",
frames[0]
);
assert!(
frames[0].get("result").is_some(),
"ping must return a successful result frame, got: {}",
frames[0]
);
}
}