use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use agent_client_protocol::DynConnectTo;
use agent_client_protocol::schema::ProtocolVersion;
use agent_client_protocol::schema::v1::{
CancelRequestNotification, ContentBlock, ContentChunk, InitializeRequest, InitializeResponse,
McpRequestId, McpServer as SchemaMcpServer, McpServerAcpId, MessageMcpNotification,
MessageMcpRequest, MessageMcpResponse, NewSessionRequest, NewSessionResponse, PermissionOption,
PermissionOptionKind, PromptRequest, PromptResponse, RequestId, RequestPermissionOutcome,
RequestPermissionRequest, RequestPermissionResponse, SelectedPermissionOutcome, SessionId,
SessionNotification, SessionUpdate, StopReason, ToolCallUpdate, ToolCallUpdateFields,
};
use agent_client_protocol::{
Agent, ByteStreams, Client, Conductor, ConnectTo, ConnectionTo, Error, JsonRpcRequest,
JsonRpcResponse, NullRun, Proxy, Responder, Role, SentRequest,
mcp_server::{McpConnectionTo, McpServer, McpServerConnect},
role,
};
use agent_client_protocol_conductor::{ConductorImpl, ProxiesAndAgent};
use futures::StreamExt as _;
use futures::channel::mpsc;
use serde::{Deserialize, Serialize};
use tokio::io::duplex;
use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt};
#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)]
#[request(method = "test/simple", response = SimpleResponse)]
struct SimpleRequest {
message: String,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)]
struct SimpleResponse {
result: String,
}
#[derive(Clone)]
struct TrackingMcpServer {
connect_tx: mpsc::UnboundedSender<McpServerAcpId>,
}
impl<Counterpart: Role> McpServerConnect<Counterpart> for TrackingMcpServer {
fn name(&self) -> String {
"tracking-mcp".to_string()
}
fn connect(&self, cx: McpConnectionTo<Counterpart>) -> DynConnectTo<role::mcp::Client> {
self.connect_tx
.unbounded_send(
cx.server_id()
.expect("cancellation test server is attached through ACP")
.clone(),
)
.unwrap();
DynConnectTo::new(EmptyMcpServerComponent)
}
}
struct EmptyMcpServerComponent;
impl ConnectTo<role::mcp::Client> for EmptyMcpServerComponent {
async fn connect_to(self, client: impl ConnectTo<role::mcp::Server>) -> Result<(), Error> {
role::mcp::Server
.builder()
.connect_with(client, async |_cx| {
std::future::pending::<Result<(), Error>>().await
})
.await
}
}
async fn next_with_timeout<T>(rx: &mut mpsc::UnboundedReceiver<T>) -> T {
tokio::time::timeout(Duration::from_secs(10), rx.next())
.await
.expect("timed out waiting for channel event")
.expect("channel closed before expected event")
}
fn assert_no_event<T: std::fmt::Debug>(rx: &mut mpsc::UnboundedReceiver<T>) {
if let Ok(event) = rx.try_recv() {
panic!("unexpected event: {event:?}");
}
}
fn advertised_mcp_server_id(request: &NewSessionRequest) -> McpServerAcpId {
match request.mcp_servers.as_slice() {
[SchemaMcpServer::Acp(acp)] => acp.server_id.clone(),
servers => panic!("expected exactly one ACP MCP server, got {servers:?}"),
}
}
struct InProcessArrowProxy;
impl ConnectTo<Conductor> for InProcessArrowProxy {
async fn connect_to(self, client: impl ConnectTo<Proxy>) -> Result<(), Error> {
agent_client_protocol_test::arrow_proxy::run_arrow_proxy(client).await
}
}
fn prompt_text(request: &PromptRequest) -> String {
request
.prompt
.iter()
.filter_map(|block| match block {
ContentBlock::Text(text) => Some(text.text.as_str()),
_ => None,
})
.collect()
}
#[tokio::test]
async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Error> {
let (agent_cancel_tx, mut agent_cancel_rx) = mpsc::unbounded();
let (parked_id_tx, mut parked_id_rx) = mpsc::unbounded();
let agent = Agent
.builder()
.on_receive_request(
async |initialize: InitializeRequest, responder, _cx: ConnectionTo<Client>| {
responder.respond(InitializeResponse::new(initialize.protocol_version))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_request(
async move |request: SimpleRequest,
responder: Responder<SimpleResponse>,
cx: ConnectionTo<Client>| {
if request.message == "park" {
parked_id_tx.unbounded_send(responder.id().clone()).unwrap();
let cancellation = responder.cancellation();
cx.spawn(async move {
let response = cancellation
.run_until_cancelled(std::future::pending::<
Result<SimpleResponse, Error>,
>())
.await;
responder.respond_with_result(response)
})?;
return Ok(());
}
responder.respond(SimpleResponse {
result: format!("echo: {}", request.message),
})
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_notification(
async move |cancel: CancelRequestNotification, _cx: ConnectionTo<Client>| {
agent_cancel_tx.unbounded_send(cancel.request_id).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
);
let (editor_write, conductor_read) = duplex(8192);
let (conductor_write, editor_read) = duplex(8192);
let conductor_handle = tokio::spawn(async move {
ConductorImpl::new_agent(
"cancellation-conductor".to_string(),
ProxiesAndAgent::new(agent).proxy(Proxy.builder()),
)
.run(ByteStreams::new(
conductor_write.compat_write(),
conductor_read.compat(),
))
.await
});
let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move {
Client
.builder()
.connect_with(
ByteStreams::new(editor_write.compat_write(), editor_read.compat()),
async |cx| {
let initialize = cx
.send_request(InitializeRequest::new(ProtocolVersion::V1))
.block_task()
.await?;
assert_eq!(initialize.protocol_version, ProtocolVersion::V1);
let request: SentRequest<SimpleResponse> = cx.send_request(SimpleRequest {
message: "park".into(),
});
let client_request_id = request.id().clone();
request.cancel()?;
let error = request
.block_task()
.await
.expect_err("request should be cancelled");
assert_eq!(i32::from(error.code), -32800);
let barrier = cx
.send_request(SimpleRequest {
message: "barrier".into(),
})
.block_task()
.await?;
assert_eq!(barrier.result, "echo: barrier");
Ok(client_request_id)
},
)
.await
})
.await
.expect("test timed out")
.expect("client failed");
let parked_id = next_with_timeout(&mut parked_id_rx).await;
assert_ne!(
parked_id, client_request_id,
"each hop must re-issue the request under its own ID"
);
let observed = next_with_timeout(&mut agent_cancel_rx).await;
assert_eq!(observed, parked_id);
assert_no_event(&mut agent_cancel_rx);
conductor_handle.abort();
Ok(())
}
#[tokio::test]
async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Error> {
let (client_cancel_tx, mut client_cancel_rx) = mpsc::unbounded();
let (parked_id_tx, mut parked_id_rx) = mpsc::unbounded();
let agent = Agent
.builder()
.on_receive_request(
async |initialize: InitializeRequest, responder, _cx: ConnectionTo<Client>| {
responder.respond(InitializeResponse::new(initialize.protocol_version))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_request(
async |request: SimpleRequest,
responder: Responder<SimpleResponse>,
cx: ConnectionTo<Client>| {
if request.message == "trigger reverse cancel" {
let connection = cx.clone();
cx.spawn(async move {
let upstream: SentRequest<SimpleResponse> =
connection.send_request(SimpleRequest {
message: "park".into(),
});
upstream.cancel()?;
let error = upstream
.block_task()
.await
.expect_err("request to the client should be cancelled");
responder.respond(SimpleResponse {
result: format!("client request error: {}", i32::from(error.code)),
})
})?;
return Ok(());
}
responder.respond(SimpleResponse {
result: format!("echo: {}", request.message),
})
},
agent_client_protocol::on_receive_request!(),
);
let (editor_write, conductor_read) = duplex(8192);
let (conductor_write, editor_read) = duplex(8192);
let conductor_handle = tokio::spawn(async move {
ConductorImpl::new_agent(
"cancellation-conductor".to_string(),
ProxiesAndAgent::new(agent),
)
.run(ByteStreams::new(
conductor_write.compat_write(),
conductor_read.compat(),
))
.await
});
tokio::time::timeout(Duration::from_secs(30), async move {
Client
.builder()
.on_receive_request(
async move |request: SimpleRequest,
responder: Responder<SimpleResponse>,
cx: ConnectionTo<Agent>| {
assert_eq!(request.message, "park");
parked_id_tx.unbounded_send(responder.id().clone()).unwrap();
let cancellation = responder.cancellation();
cx.spawn(async move {
let response = cancellation
.run_until_cancelled(std::future::pending::<
Result<SimpleResponse, Error>,
>())
.await;
responder.respond_with_result(response)
})?;
Ok(())
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_notification(
async move |cancel: CancelRequestNotification, _cx: ConnectionTo<Agent>| {
client_cancel_tx.unbounded_send(cancel.request_id).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
)
.connect_with(
ByteStreams::new(editor_write.compat_write(), editor_read.compat()),
async |cx| {
let initialize = cx
.send_request(InitializeRequest::new(ProtocolVersion::V1))
.block_task()
.await?;
assert_eq!(initialize.protocol_version, ProtocolVersion::V1);
let response = cx
.send_request(SimpleRequest {
message: "trigger reverse cancel".into(),
})
.block_task()
.await?;
assert_eq!(response.result, "client request error: -32800");
let barrier = cx
.send_request(SimpleRequest {
message: "barrier".into(),
})
.block_task()
.await?;
assert_eq!(barrier.result, "echo: barrier");
Ok(())
},
)
.await
})
.await
.expect("test timed out")
.expect("client failed");
let parked_id = next_with_timeout(&mut parked_id_rx).await;
let observed = next_with_timeout(&mut client_cancel_rx).await;
assert_eq!(observed, parked_id);
assert_no_event(&mut client_cancel_rx);
conductor_handle.abort();
Ok(())
}
#[tokio::test]
async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), Error> {
let (agent_cancel_tx, mut agent_cancel_rx) = mpsc::unbounded();
let (prompt_id_tx, mut prompt_id_rx) = mpsc::unbounded();
let (client_cancel_tx, mut client_cancel_rx) = mpsc::unbounded();
let (permission_id_tx, mut permission_id_rx) = mpsc::unbounded();
let (session_update_tx, mut session_update_rx) = mpsc::unbounded();
let agent = Agent
.builder()
.on_receive_request(
async |initialize: InitializeRequest, responder, _cx: ConnectionTo<Client>| {
responder.respond(InitializeResponse::new(initialize.protocol_version))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_request(
async |request: NewSessionRequest, responder, _cx: ConnectionTo<Client>| {
assert_eq!(request.mcp_servers, Vec::<SchemaMcpServer>::new());
responder.respond(NewSessionResponse::new(SessionId::new("test-session")))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_request(
async move |request: PromptRequest,
responder: Responder<PromptResponse>,
cx: ConnectionTo<Client>| {
let text = prompt_text(&request);
if text != "park" {
cx.send_notification(SessionNotification::new(
request.session_id,
SessionUpdate::AgentMessageChunk(ContentChunk::new(text.into())),
))?;
return responder.respond(PromptResponse::new(StopReason::EndTurn));
}
prompt_id_tx.unbounded_send(responder.id().clone()).unwrap();
let cancellation = responder.cancellation();
let connection = cx.clone();
cx.spawn(async move {
let permission: SentRequest<RequestPermissionResponse> = connection
.send_request(RequestPermissionRequest::new(
request.session_id,
ToolCallUpdate::new("tool-1", ToolCallUpdateFields::default()),
vec![PermissionOption::new(
"allow",
"Allow",
PermissionOptionKind::AllowOnce,
)],
));
cancellation.cancelled().await;
permission.cancel()?;
let permission_error = permission
.block_task()
.await
.expect_err("permission request should be cancelled");
if i32::from(permission_error.code) == -32800 {
responder.respond_with_result(Err(Error::request_cancelled()))
} else {
responder.respond_with_result(Err(
agent_client_protocol::util::internal_error(format!(
"unexpected permission error: {permission_error:?}"
)),
))
}
})?;
Ok(())
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_notification(
async move |cancel: CancelRequestNotification, _cx: ConnectionTo<Client>| {
agent_cancel_tx.unbounded_send(cancel.request_id).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
);
let (editor_write, conductor_read) = duplex(8192);
let (conductor_write, editor_read) = duplex(8192);
let conductor_handle = tokio::spawn(async move {
ConductorImpl::new_agent(
"cancellation-conductor".to_string(),
ProxiesAndAgent::new(agent).proxy(InProcessArrowProxy),
)
.run(ByteStreams::new(
conductor_write.compat_write(),
conductor_read.compat(),
))
.await
});
let client_prompt_id = tokio::time::timeout(Duration::from_secs(30), async move {
Client
.builder()
.on_receive_request(
async move |_request: RequestPermissionRequest,
responder: Responder<RequestPermissionResponse>,
cx: ConnectionTo<Agent>| {
permission_id_tx
.unbounded_send(responder.id().clone())
.unwrap();
let cancellation = responder.cancellation();
cx.spawn(async move {
let response = cancellation
.run_until_cancelled(std::future::pending::<
Result<RequestPermissionResponse, Error>,
>())
.await;
responder.respond_with_result(response)
})?;
Ok(())
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_notification(
async move |notification: SessionNotification, _cx: ConnectionTo<Agent>| {
if let SessionUpdate::AgentMessageChunk(ContentChunk {
content: ContentBlock::Text(text),
..
}) = notification.update
{
session_update_tx.unbounded_send(text.text).unwrap();
}
Ok(())
},
agent_client_protocol::on_receive_notification!(),
)
.on_receive_notification(
async move |cancel: CancelRequestNotification, _cx: ConnectionTo<Agent>| {
client_cancel_tx.unbounded_send(cancel.request_id).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
)
.connect_with(
ByteStreams::new(editor_write.compat_write(), editor_read.compat()),
async |cx| {
let initialize = cx
.send_request(InitializeRequest::new(ProtocolVersion::V1))
.block_task()
.await?;
assert_eq!(initialize.protocol_version, ProtocolVersion::V1);
let session = cx
.send_request(NewSessionRequest::new(
std::env::current_dir().map_err(Error::into_internal_error)?,
))
.block_task()
.await?;
let prompt: SentRequest<PromptResponse> = cx.send_request(PromptRequest::new(
session.session_id.clone(),
vec!["park".into()],
));
let client_prompt_id = prompt.id().clone();
prompt.cancel()?;
let error = prompt
.block_task()
.await
.expect_err("prompt should be cancelled");
assert_eq!(i32::from(error.code), -32800);
let barrier: PromptResponse = cx
.send_request(PromptRequest::new(
session.session_id.clone(),
vec!["barrier".into()],
))
.block_task()
.await?;
assert_eq!(barrier.stop_reason, StopReason::EndTurn);
Ok(client_prompt_id)
},
)
.await
})
.await
.expect("test timed out")
.expect("client failed");
let prompt_id = next_with_timeout(&mut prompt_id_rx).await;
assert_ne!(
prompt_id, client_prompt_id,
"each hop must re-issue the request under its own ID"
);
let observed = next_with_timeout(&mut agent_cancel_rx).await;
assert_eq!(observed, prompt_id);
assert_no_event(&mut agent_cancel_rx);
let permission_id = next_with_timeout(&mut permission_id_rx).await;
let observed = next_with_timeout(&mut client_cancel_rx).await;
assert_eq!(observed, permission_id);
assert_no_event(&mut client_cancel_rx);
let update = next_with_timeout(&mut session_update_rx).await;
assert_eq!(update, ">barrier");
conductor_handle.abort();
Ok(())
}
#[tokio::test]
async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error> {
let (agent_cancel_tx, mut agent_cancel_rx) = mpsc::unbounded();
let (parked_id_tx, mut parked_id_rx) = mpsc::unbounded();
let agent = Agent
.builder()
.on_receive_request(
async |initialize: InitializeRequest, responder, _cx: ConnectionTo<Client>| {
responder.respond(InitializeResponse::new(initialize.protocol_version))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_request(
async move |request: NewSessionRequest,
responder: Responder<NewSessionResponse>,
cx: ConnectionTo<Client>| {
if request.cwd.ends_with("park-session") {
parked_id_tx.unbounded_send(responder.id().clone()).unwrap();
let cancellation = responder.cancellation();
cx.spawn(async move {
let response = cancellation
.run_until_cancelled(std::future::pending::<
Result<NewSessionResponse, Error>,
>())
.await;
responder.respond_with_result(response)
})?;
return Ok(());
}
responder.respond(NewSessionResponse::new(SessionId::new("normal-session")))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_notification(
async move |cancel: CancelRequestNotification, _cx: ConnectionTo<Client>| {
agent_cancel_tx.unbounded_send(cancel.request_id).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
);
let (editor_write, conductor_read) = duplex(8192);
let (conductor_write, editor_read) = duplex(8192);
let conductor_handle = tokio::spawn(async move {
ConductorImpl::new_agent(
"cancellation-conductor".to_string(),
ProxiesAndAgent::new(agent).proxy(Proxy.builder()),
)
.run(ByteStreams::new(
conductor_write.compat_write(),
conductor_read.compat(),
))
.await
});
let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move {
Client
.builder()
.connect_with(
ByteStreams::new(editor_write.compat_write(), editor_read.compat()),
async |cx| {
let initialize = cx
.send_request(InitializeRequest::new(ProtocolVersion::V1))
.block_task()
.await?;
assert_eq!(initialize.protocol_version, ProtocolVersion::V1);
let request: SentRequest<NewSessionResponse> =
cx.send_request(NewSessionRequest::new("/park-session"));
let client_request_id = request.id().clone();
request.cancel()?;
let error = request
.block_task()
.await
.expect_err("session/new should be cancelled");
assert_eq!(i32::from(error.code), -32800);
let session = cx
.send_request(NewSessionRequest::new(
std::env::current_dir().map_err(Error::into_internal_error)?,
))
.block_task()
.await?;
assert_eq!(session.session_id, SessionId::new("normal-session"));
Ok(client_request_id)
},
)
.await
})
.await
.expect("test timed out")
.expect("client failed");
let parked_id = next_with_timeout(&mut parked_id_rx).await;
assert_ne!(
parked_id, client_request_id,
"each hop must re-issue the request under its own ID"
);
let observed = next_with_timeout(&mut agent_cancel_rx).await;
assert_eq!(observed, parked_id);
assert_no_event(&mut agent_cancel_rx);
conductor_handle.abort();
Ok(())
}
#[tokio::test]
async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), Error> {
let (agent_cancel_tx, mut agent_cancel_rx) = mpsc::unbounded();
let (parked_id_tx, mut parked_id_rx) = mpsc::unbounded();
let agent = Agent
.builder()
.on_receive_request(
async |initialize: InitializeRequest, responder, _cx: ConnectionTo<Client>| {
responder.respond(InitializeResponse::new(initialize.protocol_version))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_request(
async move |request: NewSessionRequest,
responder: Responder<NewSessionResponse>,
cx: ConnectionTo<Client>| {
if request.cwd.ends_with("park-session") {
parked_id_tx.unbounded_send(responder.id().clone()).unwrap();
let cancellation = responder.cancellation();
cx.spawn(async move {
let response = cancellation
.run_until_cancelled(std::future::pending::<
Result<NewSessionResponse, Error>,
>())
.await;
responder.respond_with_result(response)
})?;
return Ok(());
}
responder.respond(NewSessionResponse::new(SessionId::new("normal-session")))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_notification(
async move |cancel: CancelRequestNotification, _cx: ConnectionTo<Client>| {
agent_cancel_tx.unbounded_send(cancel.request_id).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
);
let proxy = Proxy.builder().on_receive_request_from(
Client,
async |request: NewSessionRequest,
responder: Responder<NewSessionResponse>,
cx: ConnectionTo<Conductor>| {
cx.build_session_from(request)
.on_proxy_session_start(responder, async |_session_id| Ok::<(), Error>(()))
},
agent_client_protocol::on_receive_request!(),
);
let (editor_write, conductor_read) = duplex(8192);
let (conductor_write, editor_read) = duplex(8192);
let conductor_handle = tokio::spawn(async move {
ConductorImpl::new_agent(
"helper-cancellation-conductor".to_string(),
ProxiesAndAgent::new(agent).proxy(proxy),
)
.run(ByteStreams::new(
conductor_write.compat_write(),
conductor_read.compat(),
))
.await
});
let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move {
Client
.builder()
.connect_with(
ByteStreams::new(editor_write.compat_write(), editor_read.compat()),
async |cx| {
let initialize = cx
.send_request(InitializeRequest::new(ProtocolVersion::V1))
.block_task()
.await?;
assert_eq!(initialize.protocol_version, ProtocolVersion::V1);
let request: SentRequest<NewSessionResponse> =
cx.send_request(NewSessionRequest::new("/park-session"));
let client_request_id = request.id().clone();
request.cancel()?;
let error = request
.block_task()
.await
.expect_err("session/new should be cancelled");
assert_eq!(i32::from(error.code), -32800);
let session = cx
.send_request(NewSessionRequest::new(
std::env::current_dir().map_err(Error::into_internal_error)?,
))
.block_task()
.await?;
assert_eq!(session.session_id, SessionId::new("normal-session"));
Ok(client_request_id)
},
)
.await
})
.await
.expect("test timed out")
.expect("client failed");
let parked_id = next_with_timeout(&mut parked_id_rx).await;
assert_ne!(
parked_id, client_request_id,
"each hop must re-issue the request under its own ID"
);
let observed = next_with_timeout(&mut agent_cancel_rx).await;
assert_eq!(observed, parked_id);
assert_no_event(&mut agent_cancel_rx);
conductor_handle.abort();
Ok(())
}
#[tokio::test]
async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() -> Result<(), Error>
{
let (agent_cancel_tx, mut agent_cancel_rx) = mpsc::unbounded();
let (parked_id_tx, mut parked_id_rx) = mpsc::unbounded();
let (mcp_connect_tx, mut mcp_connect_rx) = mpsc::unbounded();
let (probe_barrier_tx, mut probe_barrier_rx) = mpsc::unbounded();
let cancelled_mcp_server_id = Arc::new(Mutex::new(None::<McpServerAcpId>));
let agent = Agent
.builder()
.on_receive_request(
async |initialize: InitializeRequest, responder, _cx: ConnectionTo<Client>| {
responder.respond(InitializeResponse::new(initialize.protocol_version))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_request(
{
let cancelled_mcp_server_id = cancelled_mcp_server_id.clone();
let parked_id_tx = parked_id_tx.clone();
let probe_barrier_tx = probe_barrier_tx.clone();
async move |request: NewSessionRequest,
responder: Responder<NewSessionResponse>,
cx: ConnectionTo<Client>| {
let cancelled_mcp_server_id = cancelled_mcp_server_id.clone();
let parked_id_tx = parked_id_tx.clone();
let probe_barrier_tx = probe_barrier_tx.clone();
let advertised_mcp_server_id = advertised_mcp_server_id(&request);
if request.cwd.ends_with("park-session") {
*cancelled_mcp_server_id
.lock()
.expect("cancelled MCP ID mutex poisoned") =
Some(advertised_mcp_server_id);
parked_id_tx.unbounded_send(responder.id().clone()).unwrap();
let cancellation = responder.cancellation();
cx.spawn(async move {
let response = cancellation
.run_until_cancelled(std::future::pending::<
Result<NewSessionResponse, Error>,
>())
.await;
responder.respond_with_result(response)
})?;
return Ok(());
}
responder.respond(NewSessionResponse::new(SessionId::new("normal-session")))?;
let stale_server_id = cancelled_mcp_server_id
.lock()
.expect("cancelled MCP ID mutex poisoned")
.clone()
.expect("cancelled session should have advertised an MCP server");
let connection = cx.clone();
cx.spawn(async move {
connection
.send_request(
MessageMcpRequest::new(
stale_server_id,
"stale-server-probe",
"ping",
)
.params(
serde_json::Map::from_iter([(
"_meta".to_owned(),
serde_json::json!({
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}),
)]),
),
)
.on_receiving_result(async |result| {
assert!(
result.is_err(),
"a cancelled session must no longer route its MCP server"
);
Ok(())
})?;
let barrier = connection
.send_request(RequestPermissionRequest::new(
SessionId::new("normal-session"),
ToolCallUpdate::new(
"stale-mcp-probe-barrier",
ToolCallUpdateFields::default(),
),
vec![PermissionOption::new(
"allow",
"Allow",
PermissionOptionKind::AllowOnce,
)],
))
.block_task()
.await
.map(|_| ())
.map_err(|error| i32::from(error.code));
probe_barrier_tx.unbounded_send(barrier).unwrap();
Ok(())
})
}
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_notification(
async move |cancel: CancelRequestNotification, _cx: ConnectionTo<Client>| {
agent_cancel_tx.unbounded_send(cancel.request_id).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
);
let proxy = Proxy.builder().on_receive_request_from(
Client,
async move |request: NewSessionRequest,
responder: Responder<NewSessionResponse>,
cx: ConnectionTo<Conductor>| {
let mcp_server = McpServer::new(
TrackingMcpServer {
connect_tx: mcp_connect_tx.clone(),
},
NullRun,
);
cx.build_session_from(request)
.with_mcp_server(mcp_server)?
.on_proxy_session_start(responder, async |_session_id| Ok::<(), Error>(()))
},
agent_client_protocol::on_receive_request!(),
);
let (editor_write, conductor_read) = duplex(8192);
let (conductor_write, editor_read) = duplex(8192);
let conductor_handle = tokio::spawn(async move {
ConductorImpl::new_agent(
"helper-cleanup-conductor".to_string(),
ProxiesAndAgent::new(agent).proxy(proxy),
)
.run(ByteStreams::new(
conductor_write.compat_write(),
conductor_read.compat(),
))
.await
});
let client_result = tokio::time::timeout(Duration::from_secs(30), async move {
Client
.builder()
.on_receive_request(
async |request: RequestPermissionRequest,
responder: Responder<RequestPermissionResponse>,
_cx: ConnectionTo<Agent>| {
assert_eq!(request.session_id, SessionId::new("normal-session"));
responder.respond(RequestPermissionResponse::new(
RequestPermissionOutcome::Selected(SelectedPermissionOutcome::new("allow")),
))
},
agent_client_protocol::on_receive_request!(),
)
.connect_with(
ByteStreams::new(editor_write.compat_write(), editor_read.compat()),
async move |cx| {
let initialize = cx
.send_request(InitializeRequest::new(ProtocolVersion::V1))
.block_task()
.await?;
assert_eq!(initialize.protocol_version, ProtocolVersion::V1);
let request: SentRequest<NewSessionResponse> =
cx.send_request(NewSessionRequest::new("/park-session"));
let client_request_id = request.id().clone();
let parked_id = next_with_timeout(&mut parked_id_rx).await;
request.cancel()?;
let error = request
.block_task()
.await
.expect_err("session/new should be cancelled");
assert_eq!(i32::from(error.code), -32800);
let session = cx
.send_request(NewSessionRequest::new(
std::env::current_dir().map_err(Error::into_internal_error)?,
))
.block_task()
.await?;
assert_eq!(session.session_id, SessionId::new("normal-session"));
let probe_barrier = next_with_timeout(&mut probe_barrier_rx).await;
Ok((client_request_id, parked_id, probe_barrier))
},
)
.await
})
.await
.expect("test timed out")
.expect("client failed");
let (client_request_id, parked_id, probe_barrier) = client_result;
assert_ne!(
parked_id, client_request_id,
"each hop must re-issue the request under its own ID"
);
let observed = next_with_timeout(&mut agent_cancel_rx).await;
assert_eq!(observed, parked_id);
assert_no_event(&mut agent_cancel_rx);
assert_eq!(
probe_barrier,
Ok(()),
"agent-to-client barrier should succeed after stale MCP probe"
);
assert_no_event(&mut mcp_connect_rx);
conductor_handle.abort();
Ok(())
}
#[derive(Clone)]
struct ParkedMcpServer {
started_tx: mpsc::UnboundedSender<RequestId>,
stopped_tx: mpsc::UnboundedSender<RequestId>,
dropped_tx: mpsc::UnboundedSender<()>,
late_tx: mpsc::UnboundedSender<ConnectionTo<role::mcp::Client>>,
}
impl McpServerConnect<Conductor> for ParkedMcpServer {
fn name(&self) -> String {
"parked-mcp".into()
}
fn connect(&self, cx: McpConnectionTo<Conductor>) -> DynConnectTo<role::mcp::Client> {
assert_eq!(
cx.request_id().map(ToString::to_string).as_deref(),
Some("logical-mcp-request")
);
DynConnectTo::new(ParkedMcpComponent(self.clone()))
}
}
struct ParkedMcpComponent(ParkedMcpServer);
struct ProbeOnDrop<T> {
sender: mpsc::UnboundedSender<T>,
value: Option<T>,
}
impl<T> Drop for ProbeOnDrop<T> {
fn drop(&mut self) {
if let Some(value) = self.value.take() {
drop(self.sender.unbounded_send(value));
}
}
}
impl ConnectTo<role::mcp::Client> for ParkedMcpComponent {
async fn connect_to(self, client: impl ConnectTo<role::mcp::Server>) -> Result<(), Error> {
let started_tx = self.0.started_tx;
let stopped_tx = self.0.stopped_tx;
let late_tx = self.0.late_tx;
let _backend_dropped = ProbeOnDrop {
sender: self.0.dropped_tx,
value: Some(()),
};
role::mcp::Server
.builder()
.on_receive_request(
async move |_request: McpParkRequest,
responder: Responder<McpParkResponse>,
cx: ConnectionTo<role::mcp::Client>| {
let id = responder.id().clone();
let stopped = ProbeOnDrop {
sender: stopped_tx.clone(),
value: Some(id.clone()),
};
late_tx.unbounded_send(cx.clone()).unwrap();
started_tx.unbounded_send(id).unwrap();
let cancellation = responder.cancellation();
cx.spawn(async move {
let _stopped = stopped;
let result = cancellation
.run_until_cancelled(std::future::pending::<
Result<McpParkResponse, Error>,
>())
.await;
responder.respond_with_result(result)
})
},
agent_client_protocol::on_receive_request!(),
)
.connect_to(client)
.await
}
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)]
#[request(method = "_test/park", response = McpParkResponse)]
struct McpParkRequest {}
#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)]
struct McpParkResponse {}
#[derive(Debug, Clone, Serialize, Deserialize, agent_client_protocol::JsonRpcNotification)]
#[notification(method = "_test/late")]
struct LateMcpNotification {}
#[tokio::test]
async fn mcp_request_cancellation_crosses_proxy_and_tears_down_backend() -> Result<(), Error> {
let (started_tx, mut started_rx) = mpsc::unbounded();
let (stopped_tx, mut stopped_rx) = mpsc::unbounded();
let (dropped_tx, mut dropped_rx) = mpsc::unbounded();
let (late_tx, mut late_rx) = mpsc::unbounded();
let (request_id_tx, mut request_id_rx) = mpsc::unbounded();
let (result_tx, mut result_rx) = mpsc::unbounded();
let (notification_tx, mut notification_rx) = mpsc::unbounded();
let (cancel_gate_tx, cancel_gate_rx) = tokio::sync::oneshot::channel::<()>();
let cancel_gate = Arc::new(Mutex::new(Some(cancel_gate_rx)));
let agent = Agent
.builder()
.on_receive_request(
async |request: InitializeRequest, responder, _cx: ConnectionTo<Client>| {
responder.respond(InitializeResponse::new(request.protocol_version))
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_request(
async move |request: NewSessionRequest,
responder: Responder<NewSessionResponse>,
cx: ConnectionTo<Client>| {
let server_id = advertised_mcp_server_id(&request);
responder.respond(NewSessionResponse::new(SessionId::new(
"mcp-cancel-session",
)))?;
let gate = cancel_gate
.lock()
.unwrap()
.take()
.expect("one MCP operation");
let connection = cx.clone();
let request_id_tx = request_id_tx.clone();
let result_tx = result_tx.clone();
cx.spawn(async move {
let request = connection.send_request(
MessageMcpRequest::new(
server_id,
McpRequestId::new("logical-mcp-request"),
"_test/park",
)
.params(serde_json::Map::from_iter([(
"_meta".into(),
serde_json::json!({
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}),
)])),
);
request_id_tx.unbounded_send(request.id().clone()).unwrap();
gate.await.map_err(Error::into_internal_error)?;
request.cancel()?;
let result: Result<MessageMcpResponse, Error> = request.block_task().await;
result_tx
.unbounded_send(result.map(|_| ()).map_err(|error| i32::from(error.code)))
.unwrap();
Ok(())
})
},
agent_client_protocol::on_receive_request!(),
)
.on_receive_notification(
async move |notification: MessageMcpNotification, _cx: ConnectionTo<Client>| {
notification_tx.unbounded_send(notification).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
);
let proxy = Proxy.builder().with_mcp_server(McpServer::new(
ParkedMcpServer {
started_tx,
stopped_tx,
dropped_tx,
late_tx,
},
NullRun,
));
let (editor_write, conductor_read) = duplex(8192);
let (conductor_write, editor_read) = duplex(8192);
let conductor_handle = tokio::spawn(async move {
ConductorImpl::new_agent(
"mcp-cancel-conductor".to_string(),
ProxiesAndAgent::new(agent).proxy(proxy),
)
.run(ByteStreams::new(
conductor_write.compat_write(),
conductor_read.compat(),
))
.await
});
tokio::time::timeout(Duration::from_secs(30), async move {
Client
.builder()
.connect_with(
ByteStreams::new(editor_write.compat_write(), editor_read.compat()),
async |cx| {
cx.send_request(InitializeRequest::new(ProtocolVersion::V1))
.block_task()
.await?;
cx.send_request(NewSessionRequest::new(
std::env::current_dir().map_err(Error::into_internal_error)?,
))
.block_task()
.await?;
let outer_id = next_with_timeout(&mut request_id_rx).await;
let backend_id = next_with_timeout(&mut started_rx).await;
assert_ne!(outer_id, backend_id, "JSON-RPC IDs must be hop-local");
assert_eq!(
backend_id,
RequestId::Str("logical-mcp-request".to_owned()),
"the inner MCP ID must survive the proxy unchanged"
);
cancel_gate_tx
.send(())
.expect("agent still waiting to cancel");
assert_eq!(next_with_timeout(&mut result_rx).await, Err(-32800));
assert_eq!(next_with_timeout(&mut stopped_rx).await, backend_id);
next_with_timeout(&mut dropped_rx).await;
let late = next_with_timeout(&mut late_rx).await;
assert!(
late.send_notification(LateMcpNotification {}).is_err(),
"a stopped backend must reject an attempted late notification"
);
assert_no_event(&mut notification_rx);
Ok(())
},
)
.await
})
.await
.expect("MCP cancellation timed out")?;
conductor_handle.abort();
Ok(())
}
#[tokio::test]
async fn initialize_cancellation_propagates_through_proxy() -> Result<(), Error> {
let (agent_cancel_tx, mut agent_cancel_rx) = mpsc::unbounded();
let (parked_id_tx, mut parked_id_rx) = mpsc::unbounded();
let parked_first = Arc::new(AtomicBool::new(false));
let agent = Agent.builder().on_receive_request(
{
let parked_first = parked_first.clone();
async move |initialize: InitializeRequest,
responder: Responder<InitializeResponse>,
cx: ConnectionTo<Client>| {
if !parked_first.swap(true, Ordering::SeqCst) {
parked_id_tx.unbounded_send(responder.id().clone()).unwrap();
let cancellation = responder.cancellation();
cx.spawn(async move {
let response = cancellation
.run_until_cancelled(std::future::pending::<
Result<InitializeResponse, Error>,
>())
.await;
responder.respond_with_result(response)
})?;
return Ok(());
}
responder.respond(InitializeResponse::new(initialize.protocol_version))
}
},
agent_client_protocol::on_receive_request!(),
);
let agent = agent.on_receive_notification(
async move |cancel: CancelRequestNotification, _cx: ConnectionTo<Client>| {
agent_cancel_tx.unbounded_send(cancel.request_id).unwrap();
Ok(())
},
agent_client_protocol::on_receive_notification!(),
);
let (editor_write, conductor_read) = duplex(8192);
let (conductor_write, editor_read) = duplex(8192);
let conductor_handle = tokio::spawn(async move {
ConductorImpl::new_agent(
"cancellation-conductor".to_string(),
ProxiesAndAgent::new(agent).proxy(Proxy.builder()),
)
.run(ByteStreams::new(
conductor_write.compat_write(),
conductor_read.compat(),
))
.await
});
let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move {
Client
.builder()
.connect_with(
ByteStreams::new(editor_write.compat_write(), editor_read.compat()),
async |cx| {
let request: SentRequest<InitializeResponse> =
cx.send_request(InitializeRequest::new(ProtocolVersion::V1));
let client_request_id = request.id().clone();
request.cancel()?;
let error = request
.block_task()
.await
.expect_err("initialize should be cancelled");
assert_eq!(i32::from(error.code), -32800);
let initialize = cx
.send_request(InitializeRequest::new(ProtocolVersion::V1))
.block_task()
.await?;
assert_eq!(initialize.protocol_version, ProtocolVersion::V1);
Ok(client_request_id)
},
)
.await
})
.await
.expect("test timed out")
.expect("client failed");
let parked_id = next_with_timeout(&mut parked_id_rx).await;
assert_ne!(
parked_id, client_request_id,
"each hop must re-issue the request under its own ID"
);
let observed = next_with_timeout(&mut agent_cancel_rx).await;
assert_eq!(observed, parked_id);
assert_no_event(&mut agent_cancel_rx);
conductor_handle.abort();
Ok(())
}