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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions crates/rmcp/src/service.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -1103,9 +1103,11 @@ where
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
if let Some(ct) = local_ct_pool.remove(id) {
ct.cancel();
}
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
let send = transport.send(m);
let current_span = tracing::Span::current();
response_send_tasks.spawn(async move {
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -84,7 +84,7 @@ impl StreamableHttpClient for reqwest::Client {
return Err(StreamableHttpError::UnexpectedContentType(None));
}
}
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(event_stream)
}

Expand DownExpand Up@@ -223,7 +223,7 @@ impl StreamableHttpClient for reqwest::Client {
}
match content_type.as_deref() {
Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(StreamableHttpPostResponse::Sse(event_stream, session_id))
}
Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_cancelled_response.rs
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,211 @@
//! A receiver SHOULD NOT send a response for a request it has already been told
//! to cancel. This drives a real stdio server with raw JSON-RPC: the tool blocks
//! until the request is cancelled, so its result is only produced *after* the
//! cancellation — the service loop must drop it rather than write it to the wire.

use std::{collections::BTreeSet, process::Stdio, time::Duration};

use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
model::{CallToolRequestParams, CallToolResult, ContentBlock, ServerCapabilities, ServerInfo},
service::RequestContext,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader},
process::{Child, Command},
};

const HELPER_ENV: &str = "RMCP_CANCELLED_RESPONSE_HELPER";
const READ_TIMEOUT: Duration = Duration::from_secs(10);

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancelled_request_receives_no_response() -> anyhow::Result<()> {
let mut child = spawn_helper();
let mut writer = child.stdin.take().expect("helper stdin");
let stdout = child.stdout.take().expect("helper stdout");
let mut reader = BufReader::new(stdout);

send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "raw-test-client", "version": "0.0.0" }
}
}),
)
.await?;
collect_ids_until(&mut reader, 1, READ_TIMEOUT).await?;
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
)
.await?;

// Start a request that blocks until cancelled, then cancel it. Its response is
// produced only after the cancellation arrives, so it must be suppressed.
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": { "name": "wait-for-cancel", "arguments": {} }
}),
)
.await?;
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": { "requestId": 2 }
}),
)
.await?;
// A ping proves the server is alive past the cancellation, so the absence of
// an id=2 response is genuine suppression rather than a dead connection.
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "id": 3, "method": "ping" }),
)
.await?;

let seen = collect_ids_until(&mut reader, 3, READ_TIMEOUT).await?;
assert!(seen.contains(&3));
assert!(!seen.contains(&2));

drop(writer);
wait_for_child(&mut child).await;
Ok(())
}

struct WaitForCancelServer;

impl ServerHandler for WaitForCancelServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

async fn call_tool(
&self,
_request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
context.ct.cancelled().await;
Ok(CallToolResult::success(vec![ContentBlock::text(
"late response",
)]))
}
}

#[tokio::test]
async fn cancelled_response_helper() -> anyhow::Result<()> {
if std::env::var(HELPER_ENV).as_deref() != Ok("1") {
return Ok(());
}
run_helper_server().await?;
Ok(())
}

#[cfg(feature = "local")]
async fn run_helper_server() -> anyhow::Result<()> {
tokio::task::LocalSet::new()
.run_until(serve_helper_stdio())
.await
}

#[cfg(not(feature = "local"))]
async fn run_helper_server() -> anyhow::Result<()> {
serve_helper_stdio().await
}

async fn serve_helper_stdio() -> anyhow::Result<()> {
let server = WaitForCancelServer.serve(rmcp::transport::stdio()).await?;
server.waiting().await?;
Ok(())
}

fn spawn_helper() -> Child {
let exe = std::env::current_exe().expect("current test exe");
Command::new(exe)
.arg("--exact")
.arg("cancelled_response_helper")
.arg("--quiet")
.arg("--nocapture")
.arg("--test-threads")
.arg("1")
.env(HELPER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn helper")
}

async fn wait_for_child(child: &mut Child) {
let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await;
if child.id().is_some() {
let _ = child.kill().await;
}
}

async fn send_json<W>(writer: &mut W, message: &Value) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let serialized = serde_json::to_string(message)?;
writer.write_all(serialized.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}

/// Read response lines, collecting every message id seen, until `stop_id` is seen
/// (then a short grace read to catch any straggler) or the timeout elapses.
async fn collect_ids_until<R>(
reader: &mut BufReader<R>,
stop_id: u64,
timeout: Duration,
) -> anyhow::Result<BTreeSet<u64>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut seen = BTreeSet::new();
let mut deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let mut line = String::new();
let Ok(read_result) = tokio::time::timeout(remaining, reader.read_line(&mut line)).await
else {
break;
};
if read_result? == 0 {
break;
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(trimmed) else {
continue;
};
if let Some(id) = value.get("id").and_then(Value::as_u64) {
seen.insert(id);
if id == stop_id {
// Give any late (incorrectly-sent) response a brief window to arrive.
deadline = tokio::time::Instant::now() + Duration::from_millis(300);
}
}
}
Ok(seen)
}
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
fix: don't respond to cancelled requests by DaleSeo · Pull Request #957 · modelcontextprotocol/rust-sdk · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions crates/rmcp/src/service.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -1103,9 +1103,11 @@ where
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
if let Some(ct) = local_ct_pool.remove(id) {
ct.cancel();
}
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
let send = transport.send(m);
let current_span = tracing::Span::current();
response_send_tasks.spawn(async move {
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -84,7 +84,7 @@ impl StreamableHttpClient for reqwest::Client {
return Err(StreamableHttpError::UnexpectedContentType(None));
}
}
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(event_stream)
}

Expand DownExpand Up@@ -223,7 +223,7 @@ impl StreamableHttpClient for reqwest::Client {
}
match content_type.as_deref() {
Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(StreamableHttpPostResponse::Sse(event_stream, session_id))
}
Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_cancelled_response.rs
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,211 @@
//! A receiver SHOULD NOT send a response for a request it has already been told
//! to cancel. This drives a real stdio server with raw JSON-RPC: the tool blocks
//! until the request is cancelled, so its result is only produced *after* the
//! cancellation — the service loop must drop it rather than write it to the wire.

use std::{collections::BTreeSet, process::Stdio, time::Duration};

use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
model::{CallToolRequestParams, CallToolResult, ContentBlock, ServerCapabilities, ServerInfo},
service::RequestContext,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader},
process::{Child, Command},
};

const HELPER_ENV: &str = "RMCP_CANCELLED_RESPONSE_HELPER";
const READ_TIMEOUT: Duration = Duration::from_secs(10);

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancelled_request_receives_no_response() -> anyhow::Result<()> {
let mut child = spawn_helper();
let mut writer = child.stdin.take().expect("helper stdin");
let stdout = child.stdout.take().expect("helper stdout");
let mut reader = BufReader::new(stdout);

send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "raw-test-client", "version": "0.0.0" }
}
}),
)
.await?;
collect_ids_until(&mut reader, 1, READ_TIMEOUT).await?;
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
)
.await?;

// Start a request that blocks until cancelled, then cancel it. Its response is
// produced only after the cancellation arrives, so it must be suppressed.
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": { "name": "wait-for-cancel", "arguments": {} }
}),
)
.await?;
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": { "requestId": 2 }
}),
)
.await?;
// A ping proves the server is alive past the cancellation, so the absence of
// an id=2 response is genuine suppression rather than a dead connection.
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "id": 3, "method": "ping" }),
)
.await?;

let seen = collect_ids_until(&mut reader, 3, READ_TIMEOUT).await?;
assert!(seen.contains(&3));
assert!(!seen.contains(&2));

drop(writer);
wait_for_child(&mut child).await;
Ok(())
}

struct WaitForCancelServer;

impl ServerHandler for WaitForCancelServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

async fn call_tool(
&self,
_request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
context.ct.cancelled().await;
Ok(CallToolResult::success(vec![ContentBlock::text(
"late response",
)]))
}
}

#[tokio::test]
async fn cancelled_response_helper() -> anyhow::Result<()> {
if std::env::var(HELPER_ENV).as_deref() != Ok("1") {
return Ok(());
}
run_helper_server().await?;
Ok(())
}

#[cfg(feature = "local")]
async fn run_helper_server() -> anyhow::Result<()> {
tokio::task::LocalSet::new()
.run_until(serve_helper_stdio())
.await
}

#[cfg(not(feature = "local"))]
async fn run_helper_server() -> anyhow::Result<()> {
serve_helper_stdio().await
}

async fn serve_helper_stdio() -> anyhow::Result<()> {
let server = WaitForCancelServer.serve(rmcp::transport::stdio()).await?;
server.waiting().await?;
Ok(())
}

fn spawn_helper() -> Child {
let exe = std::env::current_exe().expect("current test exe");
Command::new(exe)
.arg("--exact")
.arg("cancelled_response_helper")
.arg("--quiet")
.arg("--nocapture")
.arg("--test-threads")
.arg("1")
.env(HELPER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn helper")
}

async fn wait_for_child(child: &mut Child) {
let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await;
if child.id().is_some() {
let _ = child.kill().await;
}
}

async fn send_json<W>(writer: &mut W, message: &Value) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let serialized = serde_json::to_string(message)?;
writer.write_all(serialized.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}

/// Read response lines, collecting every message id seen, until `stop_id` is seen
/// (then a short grace read to catch any straggler) or the timeout elapses.
async fn collect_ids_until<R>(
reader: &mut BufReader<R>,
stop_id: u64,
timeout: Duration,
) -> anyhow::Result<BTreeSet<u64>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut seen = BTreeSet::new();
let mut deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let mut line = String::new();
let Ok(read_result) = tokio::time::timeout(remaining, reader.read_line(&mut line)).await
else {
break;
};
if read_result? == 0 {
break;
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(trimmed) else {
continue;
};
if let Some(id) = value.get("id").and_then(Value::as_u64) {
seen.insert(id);
if id == stop_id {
// Give any late (incorrectly-sent) response a brief window to arrive.
deadline = tokio::time::Instant::now() + Duration::from_millis(300);
}
}
}
Ok(seen)
}
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' fix: don't respond to cancelled requests by DaleSeo · Pull Request #957 · modelcontextprotocol/rust-sdk · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions crates/rmcp/src/service.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -1103,9 +1103,11 @@ where
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
if let Some(ct) = local_ct_pool.remove(id) {
ct.cancel();
}
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
let send = transport.send(m);
let current_span = tracing::Span::current();
response_send_tasks.spawn(async move {
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -84,7 +84,7 @@ impl StreamableHttpClient for reqwest::Client {
return Err(StreamableHttpError::UnexpectedContentType(None));
}
}
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(event_stream)
}

Expand DownExpand Up@@ -223,7 +223,7 @@ impl StreamableHttpClient for reqwest::Client {
}
match content_type.as_deref() {
Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(StreamableHttpPostResponse::Sse(event_stream, session_id))
}
Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_cancelled_response.rs
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,211 @@
//! A receiver SHOULD NOT send a response for a request it has already been told
//! to cancel. This drives a real stdio server with raw JSON-RPC: the tool blocks
//! until the request is cancelled, so its result is only produced *after* the
//! cancellation — the service loop must drop it rather than write it to the wire.

use std::{collections::BTreeSet, process::Stdio, time::Duration};

use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
model::{CallToolRequestParams, CallToolResult, ContentBlock, ServerCapabilities, ServerInfo},
service::RequestContext,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader},
process::{Child, Command},
};

const HELPER_ENV: &str = "RMCP_CANCELLED_RESPONSE_HELPER";
const READ_TIMEOUT: Duration = Duration::from_secs(10);

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancelled_request_receives_no_response() -> anyhow::Result<()> {
let mut child = spawn_helper();
let mut writer = child.stdin.take().expect("helper stdin");
let stdout = child.stdout.take().expect("helper stdout");
let mut reader = BufReader::new(stdout);

send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "raw-test-client", "version": "0.0.0" }
}
}),
)
.await?;
collect_ids_until(&mut reader, 1, READ_TIMEOUT).await?;
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
)
.await?;

// Start a request that blocks until cancelled, then cancel it. Its response is
// produced only after the cancellation arrives, so it must be suppressed.
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": { "name": "wait-for-cancel", "arguments": {} }
}),
)
.await?;
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": { "requestId": 2 }
}),
)
.await?;
// A ping proves the server is alive past the cancellation, so the absence of
// an id=2 response is genuine suppression rather than a dead connection.
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "id": 3, "method": "ping" }),
)
.await?;

let seen = collect_ids_until(&mut reader, 3, READ_TIMEOUT).await?;
assert!(seen.contains(&3));
assert!(!seen.contains(&2));

drop(writer);
wait_for_child(&mut child).await;
Ok(())
}

struct WaitForCancelServer;

impl ServerHandler for WaitForCancelServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

async fn call_tool(
&self,
_request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
context.ct.cancelled().await;
Ok(CallToolResult::success(vec![ContentBlock::text(
"late response",
)]))
}
}

#[tokio::test]
async fn cancelled_response_helper() -> anyhow::Result<()> {
if std::env::var(HELPER_ENV).as_deref() != Ok("1") {
return Ok(());
}
run_helper_server().await?;
Ok(())
}

#[cfg(feature = "local")]
async fn run_helper_server() -> anyhow::Result<()> {
tokio::task::LocalSet::new()
.run_until(serve_helper_stdio())
.await
}

#[cfg(not(feature = "local"))]
async fn run_helper_server() -> anyhow::Result<()> {
serve_helper_stdio().await
}

async fn serve_helper_stdio() -> anyhow::Result<()> {
let server = WaitForCancelServer.serve(rmcp::transport::stdio()).await?;
server.waiting().await?;
Ok(())
}

fn spawn_helper() -> Child {
let exe = std::env::current_exe().expect("current test exe");
Command::new(exe)
.arg("--exact")
.arg("cancelled_response_helper")
.arg("--quiet")
.arg("--nocapture")
.arg("--test-threads")
.arg("1")
.env(HELPER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn helper")
}

async fn wait_for_child(child: &mut Child) {
let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await;
if child.id().is_some() {
let _ = child.kill().await;
}
}

async fn send_json<W>(writer: &mut W, message: &Value) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let serialized = serde_json::to_string(message)?;
writer.write_all(serialized.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}

/// Read response lines, collecting every message id seen, until `stop_id` is seen
/// (then a short grace read to catch any straggler) or the timeout elapses.
async fn collect_ids_until<R>(
reader: &mut BufReader<R>,
stop_id: u64,
timeout: Duration,
) -> anyhow::Result<BTreeSet<u64>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut seen = BTreeSet::new();
let mut deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let mut line = String::new();
let Ok(read_result) = tokio::time::timeout(remaining, reader.read_line(&mut line)).await
else {
break;
};
if read_result? == 0 {
break;
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(trimmed) else {
continue;
};
if let Some(id) = value.get("id").and_then(Value::as_u64) {
seen.insert(id);
if id == stop_id {
// Give any late (incorrectly-sent) response a brief window to arrive.
deadline = tokio::time::Instant::now() + Duration::from_millis(300);
}
}
}
Ok(seen)
}
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' fix: don't respond to cancelled requests by DaleSeo · Pull Request #957 · modelcontextprotocol/rust-sdk · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions crates/rmcp/src/service.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -1103,9 +1103,11 @@ where
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
if let Some(ct) = local_ct_pool.remove(id) {
ct.cancel();
}
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
let send = transport.send(m);
let current_span = tracing::Span::current();
response_send_tasks.spawn(async move {
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -84,7 +84,7 @@ impl StreamableHttpClient for reqwest::Client {
return Err(StreamableHttpError::UnexpectedContentType(None));
}
}
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(event_stream)
}

Expand DownExpand Up@@ -223,7 +223,7 @@ impl StreamableHttpClient for reqwest::Client {
}
match content_type.as_deref() {
Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(StreamableHttpPostResponse::Sse(event_stream, session_id))
}
Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_cancelled_response.rs
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,211 @@
//! A receiver SHOULD NOT send a response for a request it has already been told
//! to cancel. This drives a real stdio server with raw JSON-RPC: the tool blocks
//! until the request is cancelled, so its result is only produced *after* the
//! cancellation — the service loop must drop it rather than write it to the wire.

use std::{collections::BTreeSet, process::Stdio, time::Duration};

use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
model::{CallToolRequestParams, CallToolResult, ContentBlock, ServerCapabilities, ServerInfo},
service::RequestContext,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader},
process::{Child, Command},
};

const HELPER_ENV: &str = "RMCP_CANCELLED_RESPONSE_HELPER";
const READ_TIMEOUT: Duration = Duration::from_secs(10);

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancelled_request_receives_no_response() -> anyhow::Result<()> {
let mut child = spawn_helper();
let mut writer = child.stdin.take().expect("helper stdin");
let stdout = child.stdout.take().expect("helper stdout");
let mut reader = BufReader::new(stdout);

send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "raw-test-client", "version": "0.0.0" }
}
}),
)
.await?;
collect_ids_until(&mut reader, 1, READ_TIMEOUT).await?;
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
)
.await?;

// Start a request that blocks until cancelled, then cancel it. Its response is
// produced only after the cancellation arrives, so it must be suppressed.
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": { "name": "wait-for-cancel", "arguments": {} }
}),
)
.await?;
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": { "requestId": 2 }
}),
)
.await?;
// A ping proves the server is alive past the cancellation, so the absence of
// an id=2 response is genuine suppression rather than a dead connection.
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "id": 3, "method": "ping" }),
)
.await?;

let seen = collect_ids_until(&mut reader, 3, READ_TIMEOUT).await?;
assert!(seen.contains(&3));
assert!(!seen.contains(&2));

drop(writer);
wait_for_child(&mut child).await;
Ok(())
}

struct WaitForCancelServer;

impl ServerHandler for WaitForCancelServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

async fn call_tool(
&self,
_request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
context.ct.cancelled().await;
Ok(CallToolResult::success(vec![ContentBlock::text(
"late response",
)]))
}
}

#[tokio::test]
async fn cancelled_response_helper() -> anyhow::Result<()> {
if std::env::var(HELPER_ENV).as_deref() != Ok("1") {
return Ok(());
}
run_helper_server().await?;
Ok(())
}

#[cfg(feature = "local")]
async fn run_helper_server() -> anyhow::Result<()> {
tokio::task::LocalSet::new()
.run_until(serve_helper_stdio())
.await
}

#[cfg(not(feature = "local"))]
async fn run_helper_server() -> anyhow::Result<()> {
serve_helper_stdio().await
}

async fn serve_helper_stdio() -> anyhow::Result<()> {
let server = WaitForCancelServer.serve(rmcp::transport::stdio()).await?;
server.waiting().await?;
Ok(())
}

fn spawn_helper() -> Child {
let exe = std::env::current_exe().expect("current test exe");
Command::new(exe)
.arg("--exact")
.arg("cancelled_response_helper")
.arg("--quiet")
.arg("--nocapture")
.arg("--test-threads")
.arg("1")
.env(HELPER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn helper")
}

async fn wait_for_child(child: &mut Child) {
let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await;
if child.id().is_some() {
let _ = child.kill().await;
}
}

async fn send_json<W>(writer: &mut W, message: &Value) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let serialized = serde_json::to_string(message)?;
writer.write_all(serialized.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}

/// Read response lines, collecting every message id seen, until `stop_id` is seen
/// (then a short grace read to catch any straggler) or the timeout elapses.
async fn collect_ids_until<R>(
reader: &mut BufReader<R>,
stop_id: u64,
timeout: Duration,
) -> anyhow::Result<BTreeSet<u64>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut seen = BTreeSet::new();
let mut deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let mut line = String::new();
let Ok(read_result) = tokio::time::timeout(remaining, reader.read_line(&mut line)).await
else {
break;
};
if read_result? == 0 {
break;
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(trimmed) else {
continue;
};
if let Some(id) = value.get("id").and_then(Value::as_u64) {
seen.insert(id);
if id == stop_id {
// Give any late (incorrectly-sent) response a brief window to arrive.
deadline = tokio::time::Instant::now() + Duration::from_millis(300);
}
}
}
Ok(seen)
}
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' fix: don't respond to cancelled requests by DaleSeo · Pull Request #957 · modelcontextprotocol/rust-sdk · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions crates/rmcp/src/service.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -1103,9 +1103,11 @@ where
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
if let Some(ct) = local_ct_pool.remove(id) {
ct.cancel();
}
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
let send = transport.send(m);
let current_span = tracing::Span::current();
response_send_tasks.spawn(async move {
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -84,7 +84,7 @@ impl StreamableHttpClient for reqwest::Client {
return Err(StreamableHttpError::UnexpectedContentType(None));
}
}
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(event_stream)
}

Expand DownExpand Up@@ -223,7 +223,7 @@ impl StreamableHttpClient for reqwest::Client {
}
match content_type.as_deref() {
Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(StreamableHttpPostResponse::Sse(event_stream, session_id))
}
Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_cancelled_response.rs
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,211 @@
//! A receiver SHOULD NOT send a response for a request it has already been told
//! to cancel. This drives a real stdio server with raw JSON-RPC: the tool blocks
//! until the request is cancelled, so its result is only produced *after* the
//! cancellation — the service loop must drop it rather than write it to the wire.

use std::{collections::BTreeSet, process::Stdio, time::Duration};

use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
model::{CallToolRequestParams, CallToolResult, ContentBlock, ServerCapabilities, ServerInfo},
service::RequestContext,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader},
process::{Child, Command},
};

const HELPER_ENV: &str = "RMCP_CANCELLED_RESPONSE_HELPER";
const READ_TIMEOUT: Duration = Duration::from_secs(10);

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancelled_request_receives_no_response() -> anyhow::Result<()> {
let mut child = spawn_helper();
let mut writer = child.stdin.take().expect("helper stdin");
let stdout = child.stdout.take().expect("helper stdout");
let mut reader = BufReader::new(stdout);

send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "raw-test-client", "version": "0.0.0" }
}
}),
)
.await?;
collect_ids_until(&mut reader, 1, READ_TIMEOUT).await?;
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
)
.await?;

// Start a request that blocks until cancelled, then cancel it. Its response is
// produced only after the cancellation arrives, so it must be suppressed.
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": { "name": "wait-for-cancel", "arguments": {} }
}),
)
.await?;
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": { "requestId": 2 }
}),
)
.await?;
// A ping proves the server is alive past the cancellation, so the absence of
// an id=2 response is genuine suppression rather than a dead connection.
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "id": 3, "method": "ping" }),
)
.await?;

let seen = collect_ids_until(&mut reader, 3, READ_TIMEOUT).await?;
assert!(seen.contains(&3));
assert!(!seen.contains(&2));

drop(writer);
wait_for_child(&mut child).await;
Ok(())
}

struct WaitForCancelServer;

impl ServerHandler for WaitForCancelServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

async fn call_tool(
&self,
_request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
context.ct.cancelled().await;
Ok(CallToolResult::success(vec![ContentBlock::text(
"late response",
)]))
}
}

#[tokio::test]
async fn cancelled_response_helper() -> anyhow::Result<()> {
if std::env::var(HELPER_ENV).as_deref() != Ok("1") {
return Ok(());
}
run_helper_server().await?;
Ok(())
}

#[cfg(feature = "local")]
async fn run_helper_server() -> anyhow::Result<()> {
tokio::task::LocalSet::new()
.run_until(serve_helper_stdio())
.await
}

#[cfg(not(feature = "local"))]
async fn run_helper_server() -> anyhow::Result<()> {
serve_helper_stdio().await
}

async fn serve_helper_stdio() -> anyhow::Result<()> {
let server = WaitForCancelServer.serve(rmcp::transport::stdio()).await?;
server.waiting().await?;
Ok(())
}

fn spawn_helper() -> Child {
let exe = std::env::current_exe().expect("current test exe");
Command::new(exe)
.arg("--exact")
.arg("cancelled_response_helper")
.arg("--quiet")
.arg("--nocapture")
.arg("--test-threads")
.arg("1")
.env(HELPER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn helper")
}

async fn wait_for_child(child: &mut Child) {
let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await;
if child.id().is_some() {
let _ = child.kill().await;
}
}

async fn send_json<W>(writer: &mut W, message: &Value) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let serialized = serde_json::to_string(message)?;
writer.write_all(serialized.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}

/// Read response lines, collecting every message id seen, until `stop_id` is seen
/// (then a short grace read to catch any straggler) or the timeout elapses.
async fn collect_ids_until<R>(
reader: &mut BufReader<R>,
stop_id: u64,
timeout: Duration,
) -> anyhow::Result<BTreeSet<u64>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut seen = BTreeSet::new();
let mut deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let mut line = String::new();
let Ok(read_result) = tokio::time::timeout(remaining, reader.read_line(&mut line)).await
else {
break;
};
if read_result? == 0 {
break;
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(trimmed) else {
continue;
};
if let Some(id) = value.get("id").and_then(Value::as_u64) {
seen.insert(id);
if id == stop_id {
// Give any late (incorrectly-sent) response a brief window to arrive.
deadline = tokio::time::Instant::now() + Duration::from_millis(300);
}
}
}
Ok(seen)
}
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' fix: don't respond to cancelled requests by DaleSeo · Pull Request #957 · modelcontextprotocol/rust-sdk · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions crates/rmcp/src/service.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -1103,9 +1103,11 @@ where
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
if let Some(ct) = local_ct_pool.remove(id) {
ct.cancel();
}
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
let send = transport.send(m);
let current_span = tracing::Span::current();
response_send_tasks.spawn(async move {
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -84,7 +84,7 @@ impl StreamableHttpClient for reqwest::Client {
return Err(StreamableHttpError::UnexpectedContentType(None));
}
}
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(event_stream)
}

Expand DownExpand Up@@ -223,7 +223,7 @@ impl StreamableHttpClient for reqwest::Client {
}
match content_type.as_deref() {
Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(StreamableHttpPostResponse::Sse(event_stream, session_id))
}
Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_cancelled_response.rs
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,211 @@
//! A receiver SHOULD NOT send a response for a request it has already been told
//! to cancel. This drives a real stdio server with raw JSON-RPC: the tool blocks
//! until the request is cancelled, so its result is only produced *after* the
//! cancellation — the service loop must drop it rather than write it to the wire.

use std::{collections::BTreeSet, process::Stdio, time::Duration};

use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
model::{CallToolRequestParams, CallToolResult, ContentBlock, ServerCapabilities, ServerInfo},
service::RequestContext,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader},
process::{Child, Command},
};

const HELPER_ENV: &str = "RMCP_CANCELLED_RESPONSE_HELPER";
const READ_TIMEOUT: Duration = Duration::from_secs(10);

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancelled_request_receives_no_response() -> anyhow::Result<()> {
let mut child = spawn_helper();
let mut writer = child.stdin.take().expect("helper stdin");
let stdout = child.stdout.take().expect("helper stdout");
let mut reader = BufReader::new(stdout);

send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "raw-test-client", "version": "0.0.0" }
}
}),
)
.await?;
collect_ids_until(&mut reader, 1, READ_TIMEOUT).await?;
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
)
.await?;

// Start a request that blocks until cancelled, then cancel it. Its response is
// produced only after the cancellation arrives, so it must be suppressed.
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": { "name": "wait-for-cancel", "arguments": {} }
}),
)
.await?;
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": { "requestId": 2 }
}),
)
.await?;
// A ping proves the server is alive past the cancellation, so the absence of
// an id=2 response is genuine suppression rather than a dead connection.
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "id": 3, "method": "ping" }),
)
.await?;

let seen = collect_ids_until(&mut reader, 3, READ_TIMEOUT).await?;
assert!(seen.contains(&3));
assert!(!seen.contains(&2));

drop(writer);
wait_for_child(&mut child).await;
Ok(())
}

struct WaitForCancelServer;

impl ServerHandler for WaitForCancelServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

async fn call_tool(
&self,
_request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
context.ct.cancelled().await;
Ok(CallToolResult::success(vec![ContentBlock::text(
"late response",
)]))
}
}

#[tokio::test]
async fn cancelled_response_helper() -> anyhow::Result<()> {
if std::env::var(HELPER_ENV).as_deref() != Ok("1") {
return Ok(());
}
run_helper_server().await?;
Ok(())
}

#[cfg(feature = "local")]
async fn run_helper_server() -> anyhow::Result<()> {
tokio::task::LocalSet::new()
.run_until(serve_helper_stdio())
.await
}

#[cfg(not(feature = "local"))]
async fn run_helper_server() -> anyhow::Result<()> {
serve_helper_stdio().await
}

async fn serve_helper_stdio() -> anyhow::Result<()> {
let server = WaitForCancelServer.serve(rmcp::transport::stdio()).await?;
server.waiting().await?;
Ok(())
}

fn spawn_helper() -> Child {
let exe = std::env::current_exe().expect("current test exe");
Command::new(exe)
.arg("--exact")
.arg("cancelled_response_helper")
.arg("--quiet")
.arg("--nocapture")
.arg("--test-threads")
.arg("1")
.env(HELPER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn helper")
}

async fn wait_for_child(child: &mut Child) {
let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await;
if child.id().is_some() {
let _ = child.kill().await;
}
}

async fn send_json<W>(writer: &mut W, message: &Value) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let serialized = serde_json::to_string(message)?;
writer.write_all(serialized.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}

/// Read response lines, collecting every message id seen, until `stop_id` is seen
/// (then a short grace read to catch any straggler) or the timeout elapses.
async fn collect_ids_until<R>(
reader: &mut BufReader<R>,
stop_id: u64,
timeout: Duration,
) -> anyhow::Result<BTreeSet<u64>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut seen = BTreeSet::new();
let mut deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let mut line = String::new();
let Ok(read_result) = tokio::time::timeout(remaining, reader.read_line(&mut line)).await
else {
break;
};
if read_result? == 0 {
break;
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(trimmed) else {
continue;
};
if let Some(id) = value.get("id").and_then(Value::as_u64) {
seen.insert(id);
if id == stop_id {
// Give any late (incorrectly-sent) response a brief window to arrive.
deadline = tokio::time::Instant::now() + Duration::from_millis(300);
}
}
}
Ok(seen)
}
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' fix: don't respond to cancelled requests by DaleSeo · Pull Request #957 · modelcontextprotocol/rust-sdk · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions crates/rmcp/src/service.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -1103,9 +1103,11 @@ where
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
if let Some(ct) = local_ct_pool.remove(id) {
ct.cancel();
}
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
let send = transport.send(m);
let current_span = tracing::Span::current();
response_send_tasks.spawn(async move {
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -84,7 +84,7 @@ impl StreamableHttpClient for reqwest::Client {
return Err(StreamableHttpError::UnexpectedContentType(None));
}
}
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(event_stream)
}

Expand DownExpand Up@@ -223,7 +223,7 @@ impl StreamableHttpClient for reqwest::Client {
}
match content_type.as_deref() {
Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(StreamableHttpPostResponse::Sse(event_stream, session_id))
}
Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_cancelled_response.rs
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,211 @@
//! A receiver SHOULD NOT send a response for a request it has already been told
//! to cancel. This drives a real stdio server with raw JSON-RPC: the tool blocks
//! until the request is cancelled, so its result is only produced *after* the
//! cancellation — the service loop must drop it rather than write it to the wire.

use std::{collections::BTreeSet, process::Stdio, time::Duration};

use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
model::{CallToolRequestParams, CallToolResult, ContentBlock, ServerCapabilities, ServerInfo},
service::RequestContext,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader},
process::{Child, Command},
};

const HELPER_ENV: &str = "RMCP_CANCELLED_RESPONSE_HELPER";
const READ_TIMEOUT: Duration = Duration::from_secs(10);

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancelled_request_receives_no_response() -> anyhow::Result<()> {
let mut child = spawn_helper();
let mut writer = child.stdin.take().expect("helper stdin");
let stdout = child.stdout.take().expect("helper stdout");
let mut reader = BufReader::new(stdout);

send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "raw-test-client", "version": "0.0.0" }
}
}),
)
.await?;
collect_ids_until(&mut reader, 1, READ_TIMEOUT).await?;
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
)
.await?;

// Start a request that blocks until cancelled, then cancel it. Its response is
// produced only after the cancellation arrives, so it must be suppressed.
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": { "name": "wait-for-cancel", "arguments": {} }
}),
)
.await?;
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": { "requestId": 2 }
}),
)
.await?;
// A ping proves the server is alive past the cancellation, so the absence of
// an id=2 response is genuine suppression rather than a dead connection.
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "id": 3, "method": "ping" }),
)
.await?;

let seen = collect_ids_until(&mut reader, 3, READ_TIMEOUT).await?;
assert!(seen.contains(&3));
assert!(!seen.contains(&2));

drop(writer);
wait_for_child(&mut child).await;
Ok(())
}

struct WaitForCancelServer;

impl ServerHandler for WaitForCancelServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

async fn call_tool(
&self,
_request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
context.ct.cancelled().await;
Ok(CallToolResult::success(vec![ContentBlock::text(
"late response",
)]))
}
}

#[tokio::test]
async fn cancelled_response_helper() -> anyhow::Result<()> {
if std::env::var(HELPER_ENV).as_deref() != Ok("1") {
return Ok(());
}
run_helper_server().await?;
Ok(())
}

#[cfg(feature = "local")]
async fn run_helper_server() -> anyhow::Result<()> {
tokio::task::LocalSet::new()
.run_until(serve_helper_stdio())
.await
}

#[cfg(not(feature = "local"))]
async fn run_helper_server() -> anyhow::Result<()> {
serve_helper_stdio().await
}

async fn serve_helper_stdio() -> anyhow::Result<()> {
let server = WaitForCancelServer.serve(rmcp::transport::stdio()).await?;
server.waiting().await?;
Ok(())
}

fn spawn_helper() -> Child {
let exe = std::env::current_exe().expect("current test exe");
Command::new(exe)
.arg("--exact")
.arg("cancelled_response_helper")
.arg("--quiet")
.arg("--nocapture")
.arg("--test-threads")
.arg("1")
.env(HELPER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn helper")
}

async fn wait_for_child(child: &mut Child) {
let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await;
if child.id().is_some() {
let _ = child.kill().await;
}
}

async fn send_json<W>(writer: &mut W, message: &Value) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let serialized = serde_json::to_string(message)?;
writer.write_all(serialized.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}

/// Read response lines, collecting every message id seen, until `stop_id` is seen
/// (then a short grace read to catch any straggler) or the timeout elapses.
async fn collect_ids_until<R>(
reader: &mut BufReader<R>,
stop_id: u64,
timeout: Duration,
) -> anyhow::Result<BTreeSet<u64>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut seen = BTreeSet::new();
let mut deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let mut line = String::new();
let Ok(read_result) = tokio::time::timeout(remaining, reader.read_line(&mut line)).await
else {
break;
};
if read_result? == 0 {
break;
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(trimmed) else {
continue;
};
if let Some(id) = value.get("id").and_then(Value::as_u64) {
seen.insert(id);
if id == stop_id {
// Give any late (incorrectly-sent) response a brief window to arrive.
deadline = tokio::time::Instant::now() + Duration::from_millis(300);
}
}
}
Ok(seen)
}
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })(); fix: don't respond to cancelled requests by DaleSeo · Pull Request #957 · modelcontextprotocol/rust-sdk · GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions crates/rmcp/src/service.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -1103,9 +1103,11 @@ where
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
if let Some(ct) = local_ct_pool.remove(id) {
ct.cancel();
}
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
let send = transport.send(m);
let current_span = tracing::Span::current();
response_send_tasks.spawn(async move {
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -84,7 +84,7 @@ impl StreamableHttpClient for reqwest::Client {
return Err(StreamableHttpError::UnexpectedContentType(None));
}
}
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(event_stream)
}

Expand DownExpand Up@@ -223,7 +223,7 @@ impl StreamableHttpClient for reqwest::Client {
}
match content_type.as_deref() {
Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
let event_stream = SseStream::from_byte_stream(response.bytes_stream()).boxed();
let event_stream = SseStream::from_bytes_stream(response.bytes_stream()).boxed();
Ok(StreamableHttpPostResponse::Sse(event_stream, session_id))
}
Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
Expand Down
211 changes: 211 additions & 0 deletions crates/rmcp/tests/test_cancelled_response.rs
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,211 @@
//! A receiver SHOULD NOT send a response for a request it has already been told
//! to cancel. This drives a real stdio server with raw JSON-RPC: the tool blocks
//! until the request is cancelled, so its result is only produced *after* the
//! cancellation — the service loop must drop it rather than write it to the wire.

use std::{collections::BTreeSet, process::Stdio, time::Duration};

use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
model::{CallToolRequestParams, CallToolResult, ContentBlock, ServerCapabilities, ServerInfo},
service::RequestContext,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader},
process::{Child, Command},
};

const HELPER_ENV: &str = "RMCP_CANCELLED_RESPONSE_HELPER";
const READ_TIMEOUT: Duration = Duration::from_secs(10);

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancelled_request_receives_no_response() -> anyhow::Result<()> {
let mut child = spawn_helper();
let mut writer = child.stdin.take().expect("helper stdin");
let stdout = child.stdout.take().expect("helper stdout");
let mut reader = BufReader::new(stdout);

send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "raw-test-client", "version": "0.0.0" }
}
}),
)
.await?;
collect_ids_until(&mut reader, 1, READ_TIMEOUT).await?;
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }),
)
.await?;

// Start a request that blocks until cancelled, then cancel it. Its response is
// produced only after the cancellation arrives, so it must be suppressed.
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": { "name": "wait-for-cancel", "arguments": {} }
}),
)
.await?;
send_json(
&mut writer,
&json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": { "requestId": 2 }
}),
)
.await?;
// A ping proves the server is alive past the cancellation, so the absence of
// an id=2 response is genuine suppression rather than a dead connection.
send_json(
&mut writer,
&json!({ "jsonrpc": "2.0", "id": 3, "method": "ping" }),
)
.await?;

let seen = collect_ids_until(&mut reader, 3, READ_TIMEOUT).await?;
assert!(seen.contains(&3));
assert!(!seen.contains(&2));

drop(writer);
wait_for_child(&mut child).await;
Ok(())
}

struct WaitForCancelServer;

impl ServerHandler for WaitForCancelServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}

async fn call_tool(
&self,
_request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
context.ct.cancelled().await;
Ok(CallToolResult::success(vec![ContentBlock::text(
"late response",
)]))
}
}

#[tokio::test]
async fn cancelled_response_helper() -> anyhow::Result<()> {
if std::env::var(HELPER_ENV).as_deref() != Ok("1") {
return Ok(());
}
run_helper_server().await?;
Ok(())
}

#[cfg(feature = "local")]
async fn run_helper_server() -> anyhow::Result<()> {
tokio::task::LocalSet::new()
.run_until(serve_helper_stdio())
.await
}

#[cfg(not(feature = "local"))]
async fn run_helper_server() -> anyhow::Result<()> {
serve_helper_stdio().await
}

async fn serve_helper_stdio() -> anyhow::Result<()> {
let server = WaitForCancelServer.serve(rmcp::transport::stdio()).await?;
server.waiting().await?;
Ok(())
}

fn spawn_helper() -> Child {
let exe = std::env::current_exe().expect("current test exe");
Command::new(exe)
.arg("--exact")
.arg("cancelled_response_helper")
.arg("--quiet")
.arg("--nocapture")
.arg("--test-threads")
.arg("1")
.env(HELPER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn helper")
}

async fn wait_for_child(child: &mut Child) {
let _ = tokio::time::timeout(Duration::from_secs(2), child.wait()).await;
if child.id().is_some() {
let _ = child.kill().await;
}
}

async fn send_json<W>(writer: &mut W, message: &Value) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let serialized = serde_json::to_string(message)?;
writer.write_all(serialized.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}

/// Read response lines, collecting every message id seen, until `stop_id` is seen
/// (then a short grace read to catch any straggler) or the timeout elapses.
async fn collect_ids_until<R>(
reader: &mut BufReader<R>,
stop_id: u64,
timeout: Duration,
) -> anyhow::Result<BTreeSet<u64>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut seen = BTreeSet::new();
let mut deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
break;
}
let mut line = String::new();
let Ok(read_result) = tokio::time::timeout(remaining, reader.read_line(&mut line)).await
else {
break;
};
if read_result? == 0 {
break;
}
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(trimmed) else {
continue;
};
if let Some(id) = value.get("id").and_then(Value::as_u64) {
seen.insert(id);
if id == stop_id {
// Give any late (incorrectly-sent) response a brief window to arrive.
deadline = tokio::time::Instant::now() + Duration::from_millis(300);
}
}
}
Ok(seen)
}