From ee9bf7ed5207e4f247931101ef7ae0b216241183 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Thu, 24 Sep 2026 13:51:33 +0200 Subject: [PATCH 1/8] feat(acp): implement request-scoped native MCP transport --- Cargo.lock | 3 +- Cargo.toml | 3 +- src/agent-client-protocol/src/jsonrpc.rs | 2 +- .../src/mcp_server/active_session.rs | 659 +++++++----------- .../src/mcp_server/context.rs | 30 +- .../src/schema/enum_impls.rs | 14 - src/agent-client-protocol/src/schema/mcp.rs | 11 +- .../src/schema/v2_impls.rs | 22 - .../tests/meta_propagation.rs | 5 +- .../tests/protocol_v2.rs | 89 +-- .../tests/session_v2_mcp.rs | 93 +-- 11 files changed, 321 insertions(+), 610 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 4ce55f3b..19843773 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -140,8 +140,7 @@ dependencies = [ [[package]] name = "agent-client-protocol-schema" version = "1.9.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1d595a79e1665d02c91c7dd1bedd118501e75cc0a3d5b1651b5150ec4c5c1e0d" +source = "git+https://github.com/agentclientprotocol/agent-client-protocol?rev=1ae7f09519fa0ba43289365da42bd589468e135f#1ae7f09519fa0ba43289365da42bd589468e135f" dependencies = [ "anyhow", "derive_more", diff --git a/Cargo.toml b/Cargo.toml index e62bf28d..4c55ba77 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -35,7 +35,8 @@ agent-client-protocol-trace-viewer = { path = "src/agent-client-protocol-trace-v yopo = { package = "agent-client-protocol-yopo", path = "src/yopo" } # Protocol -agent-client-protocol-schema = { version = "=1.9.1", default-features = false, features = ["tracing"] } +# Draft cross-repository validation; replace with the released schema before publishing. +agent-client-protocol-schema = { git = "https://github.com/agentclientprotocol/agent-client-protocol", rev = "1ae7f09519fa0ba43289365da42bd589468e135f", default-features = false, features = ["tracing"] } # Core async runtime tokio = { version = "1.52", default-features = false } diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 699ae9cf..bae99ca1 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -236,7 +236,7 @@ impl Serialize for TransportBatch { } impl TransportFrame { - fn inspect_messages( + pub(crate) fn inspect_messages( &self, observer: &mut impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error>, ) -> Result<(), crate::Error> { diff --git a/src/agent-client-protocol/src/mcp_server/active_session.rs b/src/agent-client-protocol/src/mcp_server/active_session.rs index c30322cf..a20d23d5 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -1,213 +1,122 @@ -use std::{marker::PhantomData, sync::Arc}; +//! Request-scoped native MCP transport. An ACP request owns exactly one backend instance. -use futures::channel::mpsc; -use futures::{SinkExt, StreamExt}; -use rustc_hash::FxHashMap; +use futures::{ + StreamExt, + channel::oneshot, + future::{self, Either}, +}; use serde_json::{Map, Value}; - -use crate::mcp_server::{McpConnectionContext, McpConnectionTo, McpServerConnect}; -use crate::role; -use crate::role::HasPeer; -use crate::schema::v1::{ - ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, DisconnectMcpResponse, - McpConnectionId, McpServerAcpId, MessageMcpNotification, MessageMcpRequest, MessageMcpResponse, +use std::{ + collections::HashMap, + marker::PhantomData, + sync::{Arc, Mutex, Weak}, }; -use crate::util::MatchDispatchFrom; + use crate::{ Agent, Channel, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, Responder, Role, UntypedMessage, + JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, RawJsonRpcMessage, Responder, Role, + TransportFrame, + mcp_server::{McpConnectionContext, McpConnectionTo, McpServerConnect}, + role::HasPeer, + schema::v1::{ + McpRequestId, McpServerAcpId, MessageMcpNotification, MessageMcpRequest, + MessageMcpResponse, RequestId, + }, + util::MatchDispatchFrom, }; -/// Stable protocol v1 native MCP-over-ACP wire types. pub(super) struct V1McpProtocol; - -/// Draft protocol v2 native MCP-over-ACP wire types. #[cfg(feature = "unstable_protocol_v2")] pub(super) struct V2McpProtocol; -pub(super) struct McpMessage { - method: String, - params: Option>, -} - pub(super) trait McpProtocol: Send + 'static { - type ConnectRequest: JsonRpcRequest; - type ConnectResponse: JsonRpcResponse; type MessageRequest: JsonRpcRequest; - type MessageNotification: JsonRpcNotification; type MessageResponse: JsonRpcResponse; - type DisconnectRequest: JsonRpcRequest; - type DisconnectResponse: JsonRpcResponse; + type MessageNotification: JsonRpcNotification; - fn connect_server_id(request: &Self::ConnectRequest) -> McpServerAcpId; - fn connect_response(connection_id: McpConnectionId) -> Self::ConnectResponse; - fn message_request( - connection_id: McpConnectionId, - method: String, - params: Option>, - ) -> Self::MessageRequest; - fn message_notification( - connection_id: McpConnectionId, + fn server_id(request: &Self::MessageRequest) -> McpServerAcpId; + fn request_id(request: &Self::MessageRequest) -> McpRequestId; + fn into_request(request: Self::MessageRequest) -> (String, Option>); + fn notification( + server_id: McpServerAcpId, + request_id: McpRequestId, method: String, params: Option>, ) -> Self::MessageNotification; - fn message_request_connection_id(request: &Self::MessageRequest) -> McpConnectionId; - fn message_notification_connection_id( - notification: &Self::MessageNotification, - ) -> McpConnectionId; - fn into_message_request(request: Self::MessageRequest) -> McpMessage; - fn into_message_notification(notification: Self::MessageNotification) -> McpMessage; - fn disconnect_connection_id(request: &Self::DisconnectRequest) -> McpConnectionId; - fn disconnect_response() -> Self::DisconnectResponse; } impl McpProtocol for V1McpProtocol { - type ConnectRequest = ConnectMcpRequest; - type ConnectResponse = ConnectMcpResponse; type MessageRequest = MessageMcpRequest; - type MessageNotification = MessageMcpNotification; type MessageResponse = MessageMcpResponse; - type DisconnectRequest = DisconnectMcpRequest; - type DisconnectResponse = DisconnectMcpResponse; + type MessageNotification = MessageMcpNotification; - fn connect_server_id(request: &Self::ConnectRequest) -> McpServerAcpId { + fn server_id(request: &Self::MessageRequest) -> McpServerAcpId { request.server_id.clone() } - - fn connect_response(connection_id: McpConnectionId) -> Self::ConnectResponse { - ConnectMcpResponse::new(connection_id) + fn request_id(request: &Self::MessageRequest) -> McpRequestId { + request.request_id.clone() } - - fn message_request( - connection_id: McpConnectionId, - method: String, - params: Option>, - ) -> Self::MessageRequest { - MessageMcpRequest::new(connection_id, method).params(params) + fn into_request(request: Self::MessageRequest) -> (String, Option>) { + (request.method, request.params) } - - fn message_notification( - connection_id: McpConnectionId, + fn notification( + server_id: McpServerAcpId, + request_id: McpRequestId, method: String, params: Option>, ) -> Self::MessageNotification { - MessageMcpNotification::new(connection_id, method).params(params) - } - - fn message_request_connection_id(request: &Self::MessageRequest) -> McpConnectionId { - request.connection_id.clone() - } - - fn message_notification_connection_id( - notification: &Self::MessageNotification, - ) -> McpConnectionId { - notification.connection_id.clone() - } - - fn into_message_request(request: Self::MessageRequest) -> McpMessage { - McpMessage { - method: request.method, - params: request.params, - } - } - - fn into_message_notification(notification: Self::MessageNotification) -> McpMessage { - McpMessage { - method: notification.method, - params: notification.params, - } - } - - fn disconnect_connection_id(request: &Self::DisconnectRequest) -> McpConnectionId { - request.connection_id.clone() - } - - fn disconnect_response() -> Self::DisconnectResponse { - DisconnectMcpResponse::new() + MessageMcpNotification::new(server_id, request_id, method).params(params) } } #[cfg(feature = "unstable_protocol_v2")] impl McpProtocol for V2McpProtocol { - type ConnectRequest = crate::schema::v2::ConnectMcpRequest; - type ConnectResponse = crate::schema::v2::ConnectMcpResponse; type MessageRequest = crate::schema::v2::MessageMcpRequest; - type MessageNotification = crate::schema::v2::MessageMcpNotification; type MessageResponse = crate::schema::v2::MessageMcpResponse; - type DisconnectRequest = crate::schema::v2::DisconnectMcpRequest; - type DisconnectResponse = crate::schema::v2::DisconnectMcpResponse; + type MessageNotification = crate::schema::v2::MessageMcpNotification; - fn connect_server_id(request: &Self::ConnectRequest) -> McpServerAcpId { + fn server_id(request: &Self::MessageRequest) -> McpServerAcpId { McpServerAcpId::new(request.server_id.0.clone()) } - - fn connect_response(connection_id: McpConnectionId) -> Self::ConnectResponse { - crate::schema::v2::ConnectMcpResponse::new(connection_id.0) + fn request_id(request: &Self::MessageRequest) -> McpRequestId { + McpRequestId::new(request.request_id.0.clone()) } - - fn message_request( - connection_id: McpConnectionId, - method: String, - params: Option>, - ) -> Self::MessageRequest { - crate::schema::v2::MessageMcpRequest::new(connection_id.0, method).params(params) + fn into_request(request: Self::MessageRequest) -> (String, Option>) { + (request.method, request.params) } - - fn message_notification( - connection_id: McpConnectionId, + fn notification( + server_id: McpServerAcpId, + request_id: McpRequestId, method: String, params: Option>, ) -> Self::MessageNotification { - crate::schema::v2::MessageMcpNotification::new(connection_id.0, method).params(params) - } - - fn message_request_connection_id(request: &Self::MessageRequest) -> McpConnectionId { - McpConnectionId::new(request.connection_id.0.clone()) - } - - fn message_notification_connection_id( - notification: &Self::MessageNotification, - ) -> McpConnectionId { - McpConnectionId::new(notification.connection_id.0.clone()) - } - - fn into_message_request(request: Self::MessageRequest) -> McpMessage { - McpMessage { - method: request.method, - params: request.params, - } - } - - fn into_message_notification(notification: Self::MessageNotification) -> McpMessage { - McpMessage { - method: notification.method, - params: notification.params, - } - } - - fn disconnect_connection_id(request: &Self::DisconnectRequest) -> McpConnectionId { - McpConnectionId::new(request.connection_id.0.clone()) - } - - fn disconnect_response() -> Self::DisconnectResponse { - crate::schema::v2::DisconnectMcpResponse::new() + crate::schema::v2::MessageMcpNotification::new(server_id.0, request_id.0, method) + .params(params) } } -/// The message handler for an MCP server offered to a particular session. -/// This is added as a dynamic handler to the connection context and handles -/// native MCP-over-ACP messages for the declared server ID. +/// Active operations belong to the handler; dropping the declaration closes every operation. pub(super) struct McpActiveSession { - /// The opaque ACP transport identifier for this MCP server. server_id: McpServerAcpId, - - /// The MCP server we are managing. mcp_connect: Arc>, + active: Arc>>>, + protocol: PhantomData Protocol>, +} - /// Active connections to MCP server tasks. - connections: FxHashMap>, +struct ActiveRequest { + active: Weak>>>, + id: McpRequestId, +} - protocol: PhantomData Protocol>, +impl Drop for ActiveRequest { + fn drop(&mut self) { + if let Some(active) = self.active.upgrade() { + active + .lock() + .expect("MCP request registry poisoned") + .remove(&self.id); + } + } } impl McpActiveSession @@ -222,226 +131,202 @@ where Self { server_id, mcp_connect, - connections: FxHashMap::default(), + active: Arc::default(), protocol: PhantomData, } } - /// Handle a connection request for our MCP server by creating a new MCP connection. - fn handle_connect_request( + fn handle_request( &mut self, - request: Protocol::ConnectRequest, - responder: Responder, - acp_connection: &ConnectionTo, + request: Protocol::MessageRequest, + responder: Responder, + connection: &ConnectionTo, ) -> Result< Handled<( - Protocol::ConnectRequest, - Responder, + Protocol::MessageRequest, + Responder, )>, crate::Error, > { - let server_id = Protocol::connect_server_id(&request); + let server_id = Protocol::server_id(&request); if server_id != self.server_id { return Ok(Handled::No { message: (request, responder), retry: false, }); } - - let connection_id = - McpConnectionId::new(format!("mcp-over-acp-connection:{}", uuid::Uuid::new_v4())); - let (mcp_server_tx, mut mcp_server_rx) = mpsc::channel(128); - self.connections - .insert(connection_id.clone(), mcp_server_tx); - - let (client_channel, server_channel) = Channel::duplex(); - - let client_component = { - let connection_id = connection_id.clone(); - let acp_connection = acp_connection.clone(); - - role::mcp::Client - .builder() - .on_receive_dispatch( - async move |message: Dispatch, _mcp_connection| match message { - Dispatch::Request(request, responder) => { - let (method, params) = request.into_parts(); - let params = match into_native_params(params) { - Ok(params) => params, - Err(error) => return responder.respond_with_error(error), - }; - let request = - Protocol::message_request(connection_id.clone(), method, params); - let responder = responder.wrap_params(|method, result| { - result.and_then(|response: Protocol::MessageResponse| { - response.into_json(method) - }) - }); - let message: Dispatch< - Protocol::MessageRequest, - Protocol::MessageNotification, - > = Dispatch::Request(request, responder); - acp_connection.send_proxied_message_to(Agent, message) - } - Dispatch::Notification(notification) => { - let (method, params) = notification.into_parts(); - let params = match into_native_params(params) { - Ok(params) => params, - Err(error) => { - tracing::warn!( - ?error, - "ignoring MCP notification with positional parameters" - ); - return Ok(()); - } - }; - let notification = Protocol::message_notification( - connection_id.clone(), - method, - params, - ); - let message: Dispatch< - Protocol::MessageRequest, - Protocol::MessageNotification, - > = Dispatch::Notification(notification); - acp_connection.send_proxied_message_to(Agent, message) - } - Dispatch::Response(result, router) => router.route_with_result(result), - }, - crate::on_receive_dispatch!(), - ) - .with_spawned(move |mcp_connection| async move { - // These messages were sent by the ACP agent. Forward them to the MCP server. - while let Some(message) = mcp_server_rx.next().await { - mcp_connection.send_proxied_message_to(role::mcp::Server, message)?; - } - Ok(()) - }) + let request_id = Protocol::request_id(&request); + let (method, params) = Protocol::into_request(request); + if let Err(error) = validate_modern_request(&method, params.as_ref()) { + responder.respond_with_error(error)?; + return Ok(Handled::Yes); + } + let (stop_tx, stop_rx) = oneshot::channel(); + let duplicate = { + let mut active = self.active.lock().expect("MCP request registry poisoned"); + if active.contains_key(&request_id) { + true + } else { + active.insert(request_id.clone(), stop_tx); + false + } }; + if duplicate { + responder.respond_with_error( + crate::Error::invalid_params().data("duplicate active MCP requestId"), + )?; + return Ok(Handled::Yes); + } - let spawned_server = self.mcp_connect.connect(McpConnectionTo { + let guard = ActiveRequest { + active: Arc::downgrade(&self.active), + id: request_id.clone(), + }; + let backend = self.mcp_connect.connect(McpConnectionTo { context: McpConnectionContext::Acp { - server_id, - connection_id: connection_id.clone(), + server_id: server_id.clone(), + request_id: request_id.clone(), }, - connection: acp_connection.clone(), + connection: connection.clone(), }); - - let spawn_results = acp_connection - .spawn(async move { client_component.connect_to(client_channel).await }) - .and_then(|()| { - acp_connection.spawn(async move { spawned_server.connect_to(server_channel).await }) - }); - - match spawn_results { - Ok(()) => { - responder.respond(Protocol::connect_response(connection_id))?; - Ok(Handled::Yes) - } - Err(error) => { - self.connections.remove(&connection_id); - responder.respond_with_error(error)?; - Ok(Handled::Yes) + let connection_for_task = connection.clone(); + let cancellation = responder.cancellation(); + let (mut client, server) = Channel::duplex(); + // Dropping this sender when the request completes stops the backend even if it + // has outstanding work after emitting its final response. + let (backend_stop_tx, backend_stop_rx) = oneshot::channel::<()>(); + let spawn_result = connection.spawn(async move { + let run = backend.connect_to(server); + futures::pin_mut!(run); + let stop = backend_stop_rx; + futures::pin_mut!(stop); + match future::select(run, stop).await { + Either::Left((Err(error), _)) => { + tracing::warn!(?error, "request-scoped MCP backend failed"); + } + Either::Left((Ok(()), _)) | Either::Right((_, _)) => {} } + Ok(()) + }); + if let Err(error) = spawn_result { + drop(guard); + responder.respond_with_error(error)?; + return Ok(Handled::Yes); } - } - - /// Forward a native MCP-over-ACP request to its MCP connection. - async fn handle_mcp_over_acp_request( - &mut self, - request: Protocol::MessageRequest, - responder: Responder, - ) -> Result< - Handled<( - Protocol::MessageRequest, - Responder, - )>, - crate::Error, - > { - let connection_id = Protocol::message_request_connection_id(&request); - let Some(mcp_server_tx) = self.connections.get_mut(&connection_id) else { - return Ok(Handled::No { - message: (request, responder), - retry: false, - }); - }; - let message = Protocol::into_message_request(request); - - let untyped = UntypedMessage { - method: message.method, - params: native_params_into_value(message.params), - }; - let responder = responder.wrap_params(|method, result| { - result - .and_then(|response: Value| Protocol::MessageResponse::from_value(method, response)) + let spawn_result = connection.spawn(async move { + let inner_id = RequestId::Str(request_id.0.to_string()); + let process = async { + let raw = RawJsonRpcMessage::request( + method, + params.map_or(Value::Null, Value::Object), + inner_id.clone(), + )?; + client + .tx + .unbounded_send(TransportFrame::Single(raw)) + .map_err(crate::Error::into_internal_error)?; + while let Some(frame) = client.rx.next().await { + let mut result = None; + frame.inspect_messages(&mut |message| { + // A response ends the request, even within a batch. Notifications + // following it must not escape after the operation has completed. + if result.is_some() { + return Ok(()); + } + match message { + RawJsonRpcMessage::Response(response) => { + if message.response_id() != Some(&inner_id) { + return Err(crate::Error::invalid_params() + .data("MCP backend returned a different request ID")); + } + result = Some(match response { + crate::schema::v1::Response::Result { result, .. } => { + Ok(result.clone()) + } + crate::schema::v1::Response::Error { error, .. } => { + Err(error.clone()) + } + }); + } + RawJsonRpcMessage::Notification(notification) => { + let params = match notification.params.clone() { + Some(params) => match params.into_value() { + Value::Object(map) => Some(map), + _ => return Err(crate::Error::invalid_params().data( + "MCP backend notification parameters must be an object", + )), + }, + None => None, + }; + connection_for_task.send_notification_to( + Agent, + Protocol::notification( + server_id.clone(), + request_id.clone(), + notification.method.to_string(), + params, + ), + )?; + } + RawJsonRpcMessage::Request(_) => { + return Err(crate::Error::method_not_found() + .data("reverse MCP requests are not supported")); + } + } + Ok(()) + })?; + if let Some(response) = result { + return response; + } + } + Err(crate::util::internal_error( + "MCP backend closed without a response", + )) + }; + let result = cancellation + .run_until_cancelled(async { + let process = process; + futures::pin_mut!(process); + let stop = stop_rx; + futures::pin_mut!(stop); + match future::select(process, stop).await { + Either::Left((result, _)) => result, + Either::Right((_, _)) => Err(crate::Error::request_cancelled()), + } + }) + .await; + // No more notifications can be forwarded after `process` is dropped. + // Release the ID before publishing the final response so a caller can + // immediately reuse it for the next independent operation. + drop(backend_stop_tx); + drop(guard); + let response = match result { + Ok(value) => match Protocol::MessageResponse::from_value("mcp/message", value) { + Ok(response) => responder.respond(response), + Err(error) => responder.respond_with_error(error), + }, + Err(error) => responder.respond_with_error(error), + }; + if let Err(error) = response { + tracing::debug!(?error, "cannot send request-scoped MCP response"); + } + Ok(()) }); - mcp_server_tx - .send(Dispatch::Request(untyped, responder)) - .await - .map_err(crate::Error::into_internal_error)?; - - Ok(Handled::Yes) - } - - /// Forward a native MCP-over-ACP notification to its MCP connection. - async fn handle_mcp_over_acp_notification( - &mut self, - notification: Protocol::MessageNotification, - ) -> Result, crate::Error> { - let connection_id = Protocol::message_notification_connection_id(¬ification); - let Some(mcp_server_tx) = self.connections.get_mut(&connection_id) else { - return Ok(Handled::No { - message: notification, - retry: false, - }); - }; - let message = Protocol::into_message_notification(notification); - - let untyped = UntypedMessage { - method: message.method, - params: native_params_into_value(message.params), - }; - mcp_server_tx - .send(Dispatch::Notification(untyped)) - .await - .map_err(crate::Error::into_internal_error)?; - - Ok(Handled::Yes) - } - - /// Disconnect an active native MCP-over-ACP connection. - fn handle_mcp_disconnect_request( - &mut self, - request: Protocol::DisconnectRequest, - responder: Responder, - ) -> Result< - Handled<( - Protocol::DisconnectRequest, - Responder, - )>, - crate::Error, - > { - let connection_id = Protocol::disconnect_connection_id(&request); - if self.connections.remove(&connection_id).is_none() { - return Ok(Handled::No { - message: (request, responder), - retry: false, - }); + if let Err(error) = spawn_result { + // The dropped task also drops its responder and backend stop sender. + return Err(error); } - - responder.respond(Protocol::disconnect_response())?; Ok(Handled::Yes) } } -impl HandleDispatchFrom +impl HandleDispatchFrom for McpActiveSession where Counterpart: HasPeer, - Protocol: McpProtocol, { fn describe_chain(&self) -> impl std::fmt::Debug { - "McpServerSession" + "McpServerRequests" } async fn handle_dispatch_from( @@ -450,31 +335,10 @@ where connection: ConnectionTo, ) -> Result, crate::Error> { MatchDispatchFrom::new(message, &connection) - .if_request_from( - Agent, - async |request: Protocol::ConnectRequest, responder| { - self.handle_connect_request(request, responder, &connection) - }, - ) - .await .if_request_from( Agent, async |request: Protocol::MessageRequest, responder| { - self.handle_mcp_over_acp_request(request, responder).await - }, - ) - .await - .if_notification_from( - Agent, - async |notification: Protocol::MessageNotification| { - self.handle_mcp_over_acp_notification(notification).await - }, - ) - .await - .if_request_from( - Agent, - async |request: Protocol::DisconnectRequest, responder| { - self.handle_mcp_disconnect_request(request, responder) + self.handle_request(request, responder, &connection) }, ) .await @@ -482,44 +346,49 @@ where } } -fn into_native_params(params: Value) -> Result>, crate::Error> { - match params { - Value::Null => Ok(None), - Value::Object(params) => Ok(Some(params)), - Value::Array(_) => Err(crate::Error::invalid_params() - .data("MCP-over-ACP only supports named inner MCP parameters")), - _ => { - Err(crate::Error::invalid_params() - .data("inner MCP parameters must be an object or null")) - } +fn validate_modern_request( + method: &str, + params: Option<&Map>, +) -> Result<(), crate::Error> { + if method == "initialize" { + return Err( + crate::Error::invalid_params().data("native MCP requests do not use initialize") + ); } -} - -fn native_params_into_value(params: Option>) -> Value { - params.map_or(Value::Null, Value::Object) + let meta = params + .and_then(|params| params.get("_meta")) + .and_then(Value::as_object); + if meta + .and_then(|meta| meta.get("io.modelcontextprotocol/protocolVersion")) + .and_then(Value::as_str) + != Some("2026-07-28") + || !meta + .and_then(|meta| meta.get("io.modelcontextprotocol/clientCapabilities")) + .is_some_and(Value::is_object) + { + return Err(crate::Error::invalid_params().data("inner params._meta requires io.modelcontextprotocol/protocolVersion 2026-07-28 and io.modelcontextprotocol/clientCapabilities object")); + } + Ok(()) } #[cfg(test)] mod tests { + use super::validate_modern_request; use serde_json::json; - use super::{into_native_params, native_params_into_value}; - #[test] - fn native_mcp_params_round_trip_objects_and_null() { - let object = json!({ "name": "echo", "arguments": {} }); - let params = into_native_params(object.clone()).expect("object params should be valid"); - assert_eq!(native_params_into_value(params), object); - - let params = into_native_params(serde_json::Value::Null) - .expect("omitted params should be represented as null"); - assert_eq!(native_params_into_value(params), serde_json::Value::Null); - } - - #[test] - fn native_mcp_params_reject_positional_params() { - let error = into_native_params(json!(["positional"])) - .expect_err("native MCP-over-ACP cannot represent positional params"); - assert_eq!(error.code, crate::ErrorCode::InvalidParams); + fn only_modern_request_metadata_is_accepted() { + let modern = json!({"_meta": {"io.modelcontextprotocol/protocolVersion": "2026-07-28", "io.modelcontextprotocol/clientCapabilities": {}, "requestState": {"opaque": true}}}); + assert!(validate_modern_request("tools/list", modern.as_object()).is_ok()); + assert!(validate_modern_request("initialize", modern.as_object()).is_err()); + assert!(validate_modern_request("tools/list", json!({"_meta": {"io.modelcontextprotocol/protocolVersion": "2025-03-26", "io.modelcontextprotocol/clientCapabilities": {}}}).as_object()).is_err()); + assert!( + validate_modern_request( + "tools/list", + json!({"_meta": {"io.modelcontextprotocol/protocolVersion": "2026-07-28"}}) + .as_object() + ) + .is_err() + ); } } diff --git a/src/agent-client-protocol/src/mcp_server/context.rs b/src/agent-client-protocol/src/mcp_server/context.rs index aba4b938..517f3026 100644 --- a/src/agent-client-protocol/src/mcp_server/context.rs +++ b/src/agent-client-protocol/src/mcp_server/context.rs @@ -1,7 +1,7 @@ use crate::{ConnectionTo, role::Role}; #[cfg(feature = "unstable_mcp_over_acp")] -use crate::schema::v1::{McpConnectionId, McpServerAcpId}; +use crate::schema::v1::{McpRequestId, McpServerAcpId}; /// Describes how an MCP server connection was established. #[derive(Clone, Debug, PartialEq, Eq)] @@ -16,8 +16,8 @@ pub enum McpConnectionContext { /// The identifier advertised in the session's `McpServer::Acp` declaration. server_id: McpServerAcpId, - /// The identifier for this active `mcp/connect` connection. - connection_id: McpConnectionId, + /// The logical identifier of this independent MCP request. + request_id: McpRequestId, }, } @@ -40,15 +40,15 @@ impl McpConnectionContext { } } - /// The identifier for the active `mcp/connect` connection. + /// The logical identifier of the active MCP request. /// /// Returns `None` for a standalone MCP connection. #[cfg(feature = "unstable_mcp_over_acp")] #[must_use] - pub fn connection_id(&self) -> Option<&McpConnectionId> { + pub fn request_id(&self) -> Option<&McpRequestId> { match self { Self::Standalone => None, - Self::Acp { connection_id, .. } => Some(connection_id), + Self::Acp { request_id, .. } => Some(request_id), } } } @@ -76,13 +76,13 @@ impl McpConnectionTo { self.context.server_id() } - /// The identifier for the active `mcp/connect` connection. + /// The logical identifier of the active MCP request. /// /// Returns `None` for a standalone MCP connection. #[cfg(feature = "unstable_mcp_over_acp")] #[must_use] - pub fn connection_id(&self) -> Option<&McpConnectionId> { - self.context.connection_id() + pub fn request_id(&self) -> Option<&McpRequestId> { + self.context.request_id() } /// Borrow the host protocol connection. @@ -108,24 +108,24 @@ mod tests { #[cfg(feature = "unstable_mcp_over_acp")] { assert_eq!(context.server_id(), None); - assert_eq!(context.connection_id(), None); + assert_eq!(context.request_id(), None); } } #[cfg(feature = "unstable_mcp_over_acp")] #[test] - fn acp_context_exposes_server_and_connection_ids() { - use crate::schema::v1::{McpConnectionId, McpServerAcpId}; + fn acp_context_exposes_server_and_request_ids() { + use crate::schema::v1::{McpRequestId, McpServerAcpId}; let server_id = McpServerAcpId::new("server-id"); - let connection_id = McpConnectionId::new("connection-id"); + let request_id = McpRequestId::new("request-id"); let context = McpConnectionContext::Acp { server_id: server_id.clone(), - connection_id: connection_id.clone(), + request_id: request_id.clone(), }; assert!(!context.is_standalone()); assert_eq!(context.server_id(), Some(&server_id)); - assert_eq!(context.connection_id(), Some(&connection_id)); + assert_eq!(context.request_id(), Some(&request_id)); } } diff --git a/src/agent-client-protocol/src/schema/enum_impls.rs b/src/agent-client-protocol/src/schema/enum_impls.rs index 8a937c43..e485e868 100644 --- a/src/agent-client-protocol/src/schema/enum_impls.rs +++ b/src/agent-client-protocol/src/schema/enum_impls.rs @@ -31,8 +31,6 @@ impl_jsonrpc_request_enum!(ClientRequest { SetSessionModeRequest => "session/set_mode", SetSessionConfigOptionRequest => "session/set_config_option", PromptRequest => "session/prompt", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpRequest => "mcp/message", [ext] ExtMethodRequest, }); @@ -57,8 +55,6 @@ impl_jsonrpc_response_enum!(AgentResponse { SetSessionModeResponse => "session/set_mode", SetSessionConfigOptionResponse => "session/set_config_option", PromptResponse => "session/prompt", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpResponse => "mcp/message", [ext] ExtMethodResponse, }); @@ -84,11 +80,7 @@ impl_jsonrpc_request_enum!(AgentRequest { KillTerminalRequest => "terminal/kill", CreateElicitationRequest => "elicitation/create", #[cfg(feature = "unstable_mcp_over_acp")] - ConnectMcpRequest => "mcp/connect", - #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpRequest => "mcp/message", - #[cfg(feature = "unstable_mcp_over_acp")] - DisconnectMcpRequest => "mcp/disconnect", [ext] ExtMethodRequest, }); @@ -103,18 +95,12 @@ impl_jsonrpc_response_enum!(ClientResponse { KillTerminalResponse => "terminal/kill", CreateElicitationResponse => "elicitation/create", #[cfg(feature = "unstable_mcp_over_acp")] - ConnectMcpResponse => "mcp/connect", - #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpResponse => "mcp/message", - #[cfg(feature = "unstable_mcp_over_acp")] - DisconnectMcpResponse => "mcp/disconnect", [ext] ExtMethodResponse, }); impl_jsonrpc_notification_enum!(AgentNotification { SessionNotification => "session/update", CompleteElicitationNotification => "elicitation/complete", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpNotification => "mcp/message", [ext] ExtNotification, }); diff --git a/src/agent-client-protocol/src/schema/mcp.rs b/src/agent-client-protocol/src/schema/mcp.rs index a16464bc..e5c5032d 100644 --- a/src/agent-client-protocol/src/schema/mcp.rs +++ b/src/agent-client-protocol/src/schema/mcp.rs @@ -1,15 +1,6 @@ //! JSON-RPC implementations for the unstable native MCP-over-ACP transport. -use crate::schema::v1::{ - ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, DisconnectMcpResponse, - MessageMcpNotification, MessageMcpRequest, MessageMcpResponse, -}; +use crate::schema::v1::{MessageMcpNotification, MessageMcpRequest, MessageMcpResponse}; -impl_jsonrpc_request!(ConnectMcpRequest, ConnectMcpResponse, "mcp/connect"); impl_jsonrpc_request!(MessageMcpRequest, MessageMcpResponse, "mcp/message"); impl_jsonrpc_notification!(MessageMcpNotification, "mcp/message"); -impl_jsonrpc_request!( - DisconnectMcpRequest, - DisconnectMcpResponse, - "mcp/disconnect" -); diff --git a/src/agent-client-protocol/src/schema/v2_impls.rs b/src/agent-client-protocol/src/schema/v2_impls.rs index 8cae486b..94d849e8 100644 --- a/src/agent-client-protocol/src/schema/v2_impls.rs +++ b/src/agent-client-protocol/src/schema/v2_impls.rs @@ -281,14 +281,6 @@ impl_v2_jsonrpc_request!( v2::CreateElicitationResponse, "elicitation/create" ); -#[cfg(feature = "unstable_mcp_over_acp")] -impl_v2_jsonrpc_request!(v2::ConnectMcpRequest, v2::ConnectMcpResponse, "mcp/connect"); -#[cfg(feature = "unstable_mcp_over_acp")] -impl_v2_jsonrpc_request!( - v2::DisconnectMcpRequest, - v2::DisconnectMcpResponse, - "mcp/disconnect" -); impl_v2_jsonrpc_notification!(v2::UpdateSessionNotification, "session/update"); impl_v2_jsonrpc_notification!(v2::CompleteElicitationNotification, "elicitation/complete"); @@ -316,8 +308,6 @@ impl_v2_jsonrpc_request_enum!(v2::ClientRequest { CloseSessionRequest => "session/close", SetSessionConfigOptionRequest => "session/set_config_option", PromptRequest => "session/prompt", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpRequest => "mcp/message", [ext] ExtMethodRequest, }); @@ -340,8 +330,6 @@ impl_v2_jsonrpc_response_enum!(v2::AgentResponse { CloseSessionResponse => "session/close", SetSessionConfigOptionResponse => "session/set_config_option", PromptResponse => "session/prompt", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpResponse => "mcp/message", [ext] ExtMethodResponse, }); @@ -356,11 +344,7 @@ impl_v2_jsonrpc_request_enum!(v2::AgentRequest { RequestPermissionRequest => "session/request_permission", CreateElicitationRequest => "elicitation/create", #[cfg(feature = "unstable_mcp_over_acp")] - ConnectMcpRequest => "mcp/connect", - #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpRequest => "mcp/message", - #[cfg(feature = "unstable_mcp_over_acp")] - DisconnectMcpRequest => "mcp/disconnect", [ext] ExtMethodRequest, }); @@ -368,18 +352,12 @@ impl_v2_jsonrpc_response_enum!(v2::ClientResponse { RequestPermissionResponse => "session/request_permission", CreateElicitationResponse => "elicitation/create", #[cfg(feature = "unstable_mcp_over_acp")] - ConnectMcpResponse => "mcp/connect", - #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpResponse => "mcp/message", - #[cfg(feature = "unstable_mcp_over_acp")] - DisconnectMcpResponse => "mcp/disconnect", [ext] ExtMethodResponse, }); impl_v2_jsonrpc_notification_enum!(v2::AgentNotification { UpdateSessionNotification => "session/update", CompleteElicitationNotification => "elicitation/complete", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpNotification => "mcp/message", [ext] ExtNotification, }); diff --git a/src/agent-client-protocol/tests/meta_propagation.rs b/src/agent-client-protocol/tests/meta_propagation.rs index eef3f331..537ea09f 100644 --- a/src/agent-client-protocol/tests/meta_propagation.rs +++ b/src/agent-client-protocol/tests/meta_propagation.rs @@ -106,7 +106,7 @@ fn successor_message_accepts_legacy_meta_alias() -> Result<(), agent_client_prot fn native_mcp_over_acp_message_meta_serializes_as_reserved_meta_field() -> Result<(), agent_client_protocol::Error> { let meta = trace_context_meta(); - let message = MessageMcpRequest::new("connection-1", "tools/list") + let message = MessageMcpRequest::new("server-1", "request-1", "tools/list") .params(serde_json::Map::from_iter([( "cursor".into(), Value::String("abc".into()), @@ -116,7 +116,8 @@ fn native_mcp_over_acp_message_meta_serializes_as_reserved_meta_field() let untyped = message.to_untyped_message()?; assert_eq!(untyped.method(), "mcp/message"); - assert_eq!(untyped.params()["connectionId"], "connection-1"); + assert_eq!(untyped.params()["serverId"], "server-1"); + assert_eq!(untyped.params()["requestId"], "request-1"); assert_eq!(untyped.params()["method"], "tools/list"); assert_eq!(untyped.params()["params"]["cursor"], "abc"); assert_eq!(untyped.params()["_meta"], Value::Object(meta.clone())); diff --git a/src/agent-client-protocol/tests/protocol_v2.rs b/src/agent-client-protocol/tests/protocol_v2.rs index 52a94a4a..b6f9579a 100644 --- a/src/agent-client-protocol/tests/protocol_v2.rs +++ b/src/agent-client-protocol/tests/protocol_v2.rs @@ -972,16 +972,9 @@ fn sdk_supported_v2_method_surface_is_jsonrpc_mapped() -> Result<(), Error> { .map_err(Error::into_internal_error) } - assert_client_request!( - MessageMcpRequest, - MessageMcpResponse, - "mcp/message", - v2::MessageMcpRequest::new("connection-1", "tools/list"), - message_response()? - ); assert_v2_client_notification_mapping( "mcp/message", - v2::MessageMcpNotification::new("connection-1", "notifications/tools/list"), + v2::MessageMcpNotification::new("server-1", "request-1", "notifications/tools/list"), |notification| { matches!( notification, @@ -990,37 +983,13 @@ fn sdk_supported_v2_method_surface_is_jsonrpc_mapped() -> Result<(), Error> { }, )?; - assert_agent_request!( - ConnectMcpRequest, - ConnectMcpResponse, - "mcp/connect", - v2::ConnectMcpRequest::new("server-1"), - v2::ConnectMcpResponse::new("connection-1") - ); assert_agent_request!( MessageMcpRequest, MessageMcpResponse, "mcp/message", - v2::MessageMcpRequest::new("connection-1", "tools/list"), + v2::MessageMcpRequest::new("server-1", "request-1", "tools/list"), message_response()? ); - assert_agent_request!( - DisconnectMcpRequest, - DisconnectMcpResponse, - "mcp/disconnect", - v2::DisconnectMcpRequest::new("connection-1"), - v2::DisconnectMcpResponse::new() - ); - assert_v2_agent_notification_mapping( - "mcp/message", - v2::MessageMcpNotification::new("connection-1", "notifications/tools/list"), - |notification| { - matches!( - notification, - v2::AgentNotification::MessageMcpNotification(_) - ) - }, - )?; } let cancel_params = json_value(v2::CancelRequestNotification::new(String::from( @@ -1055,72 +1024,32 @@ fn mcp_over_acp_v1_variants_are_jsonrpc_mapped() -> Result<(), Error> { }}; } - assert_message_mapping!( - v1::ClientRequest, - "mcp/message", - json_value(v1::MessageMcpRequest::new("conn-1", "tools/list"))?, - v1::ClientRequest::MessageMcpRequest(_) - ); - assert_response_mapping!( - v1::AgentResponse, - "mcp/message", - serde_json::json!({ "tools": [] }), - v1::AgentResponse::MessageMcpResponse(_) - ); assert_message_mapping!( v1::ClientNotification, "mcp/message", json_value(v1::MessageMcpNotification::new( - "conn-1", + "server-1", + "request-1", "notifications/tools/list" ))?, v1::ClientNotification::MessageMcpNotification(_) ); - assert_message_mapping!( - v1::AgentRequest, - "mcp/connect", - json_value(v1::ConnectMcpRequest::new("server-1"))?, - v1::AgentRequest::ConnectMcpRequest(_) - ); assert_message_mapping!( v1::AgentRequest, "mcp/message", - json_value(v1::MessageMcpRequest::new("conn-1", "tools/list"))?, + json_value(v1::MessageMcpRequest::new( + "server-1", + "request-1", + "tools/list" + ))?, v1::AgentRequest::MessageMcpRequest(_) ); - assert_message_mapping!( - v1::AgentRequest, - "mcp/disconnect", - json_value(v1::DisconnectMcpRequest::new("conn-1"))?, - v1::AgentRequest::DisconnectMcpRequest(_) - ); - assert_response_mapping!( - v1::ClientResponse, - "mcp/connect", - json_value(v1::ConnectMcpResponse::new("conn-1"))?, - v1::ClientResponse::ConnectMcpResponse(_) - ); assert_response_mapping!( v1::ClientResponse, "mcp/message", serde_json::json!({ "tools": [] }), v1::ClientResponse::MessageMcpResponse(_) ); - assert_response_mapping!( - v1::ClientResponse, - "mcp/disconnect", - serde_json::json!({}), - v1::ClientResponse::DisconnectMcpResponse(_) - ); - assert_message_mapping!( - v1::AgentNotification, - "mcp/message", - json_value(v1::MessageMcpNotification::new( - "conn-1", - "notifications/tools/list" - ))?, - v1::AgentNotification::MessageMcpNotification(_) - ); Ok(()) } diff --git a/src/agent-client-protocol/tests/session_v2_mcp.rs b/src/agent-client-protocol/tests/session_v2_mcp.rs index 904e7f52..f9ea571a 100644 --- a/src/agent-client-protocol/tests/session_v2_mcp.rs +++ b/src/agent-client-protocol/tests/session_v2_mcp.rs @@ -12,8 +12,8 @@ use std::{ }; use agent_client_protocol::{ - Agent, Client, ConnectTo, ConnectionTo, DynConnectTo, Error, ErrorCode, JsonRpcNotification, - JsonRpcRequest, JsonRpcResponse, Responder, RunWithConnectionTo, V2ConnectionTo, + Agent, Client, ConnectTo, ConnectionTo, DynConnectTo, Error, ErrorCode, JsonRpcRequest, + JsonRpcResponse, Responder, RunWithConnectionTo, V2ConnectionTo, mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, role, schema::{ProtocolVersion, v2}, @@ -87,21 +87,14 @@ struct ConnectionProbeResponse { nonce: String, } -#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcNotification)] -#[notification(method = "_test/notice")] -struct NoticeNotification { - message: String, -} - #[derive(Debug, PartialEq, Eq)] struct ObservedMcpContext { server_id: String, - connection_id: String, + request_id: String, } struct EchoMcpConnect { context_tx: mpsc::UnboundedSender, - notice_tx: mpsc::UnboundedSender, runner_started: Arc, dropped_tx: Mutex>>, } @@ -135,37 +128,23 @@ impl McpServerConnect for EchoMcpConnect { .server_id() .expect("the MCP server should be attached through ACP") .to_string(), - connection_id: context - .connection_id() - .expect("an attached MCP connection should have an ID") + request_id: context + .request_id() + .expect("an attached MCP request should have an ID") .to_string(), }) .expect("MCP context receiver should remain active"); - DynConnectTo::new(EchoMcpComponent { - notice_tx: self.notice_tx.clone(), - }) + DynConnectTo::new(EchoMcpComponent) } } -struct EchoMcpComponent { - notice_tx: mpsc::UnboundedSender, -} +struct EchoMcpComponent; impl ConnectTo for EchoMcpComponent { async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { - let notice_tx = self.notice_tx; - role::mcp::Server .builder() - .on_receive_notification( - async move |notification: NoticeNotification, _connection| { - notice_tx - .unbounded_send(notification.message) - .map_err(Error::into_internal_error) - }, - agent_client_protocol::on_receive_notification!(), - ) .on_receive_request( async |request: EchoRequest, responder: Responder, _connection| { responder.respond(EchoResponse { @@ -210,8 +189,7 @@ impl RunWithConnectionTo for ProbeRunner { #[derive(Debug)] struct RoundTrip { server_id: String, - connection_id: String, - notice: String, + request_id: String, response: Value, } @@ -220,36 +198,24 @@ async fn run_mcp_round_trip( server_id: &v2::McpServerAcpId, sequence: usize, ) -> Result { - let connected = connection - .send_request(v2::ConnectMcpRequest::new(server_id.clone())) - .block_task() - .await?; - let connection_id = connected.connection_id; - let notice = format!("notice-{sequence}"); - connection.send_notification( - v2::MessageMcpNotification::new(connection_id.clone(), "_test/notice") - .params(object(json!({ "message": notice }))), - )?; - + let request_id = format!("request-{sequence}"); let message = format!("message-{sequence}"); let response = connection .send_request( - v2::MessageMcpRequest::new(connection_id.clone(), "_test/echo") - .params(object(json!({ "message": message }))), + v2::MessageMcpRequest::new(server_id.clone(), request_id.clone(), "_test/echo").params( + object(json!({ "message": message, "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } })), + ), ) .block_task() .await?; let response = serde_json::from_str(response.0.get()).map_err(Error::into_internal_error)?; - connection - .send_request(v2::DisconnectMcpRequest::new(connection_id.clone())) - .block_task() - .await?; - Ok(RoundTrip { server_id: server_id.to_string(), - connection_id: connection_id.to_string(), - notice, + request_id, response, }) } @@ -258,15 +224,12 @@ async fn assert_round_trip( sequence: usize, round_trip_rx: &mut UnboundedReceiver>, context_rx: &mut UnboundedReceiver, - notice_rx: &mut UnboundedReceiver, ) -> Result<(), Error> { let round_trip = next(round_trip_rx, "MCP round trip").await?; - let context = next(context_rx, "MCP connection context").await; - let notice = next(notice_rx, "inner MCP notification").await; + let context = next(context_rx, "MCP request context").await; assert_eq!(context.server_id, round_trip.server_id); - assert_eq!(context.connection_id, round_trip.connection_id); - assert_eq!(notice, round_trip.notice); + assert_eq!(context.request_id, round_trip.request_id); assert_eq!( round_trip.response, json!({ "echoed": format!("message-{sequence}") }) @@ -364,7 +327,6 @@ async fn v2_session_mcp_attachment_is_ready_during_setup_and_lives_for_connectio let test = async move { let (context_tx, mut context_rx) = mpsc::unbounded(); - let (notice_tx, mut notice_rx) = mpsc::unbounded(); let (connector_dropped_tx, connector_dropped_rx) = oneshot::channel(); let (runner_started_tx, runner_started_rx) = oneshot::channel(); let (runner_dropped_tx, runner_dropped_rx) = oneshot::channel(); @@ -384,7 +346,6 @@ async fn v2_session_mcp_attachment_is_ready_during_setup_and_lives_for_connectio let mcp_server = McpServer::::new( EchoMcpConnect { context_tx, - notice_tx, runner_started: runner_started.clone(), dropped_tx: Mutex::new(Some(connector_dropped_tx)), }, @@ -411,7 +372,7 @@ async fn v2_session_mcp_attachment_is_ready_during_setup_and_lives_for_connectio "the MCP runner must be first-polled before session/new is published" ); - assert_round_trip(1, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(1, &mut round_trip_rx, &mut context_rx).await?; let session = pending_session.block_task().await?.into_session(); let remaining_session = session.clone(); @@ -421,7 +382,7 @@ async fn v2_session_mcp_attachment_is_ready_during_setup_and_lives_for_connectio round_trip_trigger_tx .unbounded_send(()) .map_err(Error::into_internal_error)?; - assert_round_trip(2, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(2, &mut round_trip_rx, &mut context_rx).await?; Ok(()) }) @@ -550,7 +511,6 @@ async fn v2_fork_mcp_attachment_preserves_request_and_lives_for_connection() -> let test = async move { let (context_tx, mut context_rx) = mpsc::unbounded(); - let (notice_tx, mut notice_rx) = mpsc::unbounded(); let (connector_dropped_tx, connector_dropped_rx) = oneshot::channel(); let (runner_started_tx, runner_started_rx) = oneshot::channel(); let (runner_dropped_tx, runner_dropped_rx) = oneshot::channel(); @@ -570,7 +530,6 @@ async fn v2_fork_mcp_attachment_preserves_request_and_lives_for_connection() -> let mcp_server = McpServer::::new( EchoMcpConnect { context_tx, - notice_tx, runner_started: runner_started.clone(), dropped_tx: Mutex::new(Some(connector_dropped_tx)), }, @@ -598,7 +557,7 @@ async fn v2_fork_mcp_attachment_preserves_request_and_lives_for_connection() -> "the MCP runner must be first-polled before session/fork is published" ); - assert_round_trip(1, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(1, &mut round_trip_rx, &mut context_rx).await?; let opened = pending_session.block_task().await?; assert_eq!(opened.session().session_id(), &expected_forked_session_id); @@ -612,7 +571,7 @@ async fn v2_fork_mcp_attachment_preserves_request_and_lives_for_connection() -> round_trip_trigger_tx .unbounded_send(()) .map_err(Error::into_internal_error)?; - assert_round_trip(2, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(2, &mut round_trip_rx, &mut context_rx).await?; Ok(()) }) @@ -742,7 +701,6 @@ async fn v2_resume_mcp_attachment_preserves_request_and_lives_for_connection() - let test = async move { let (context_tx, mut context_rx) = mpsc::unbounded(); - let (notice_tx, mut notice_rx) = mpsc::unbounded(); let (connector_dropped_tx, connector_dropped_rx) = oneshot::channel(); let (runner_started_tx, runner_started_rx) = oneshot::channel(); let (runner_dropped_tx, runner_dropped_rx) = oneshot::channel(); @@ -762,7 +720,6 @@ async fn v2_resume_mcp_attachment_preserves_request_and_lives_for_connection() - let mcp_server = McpServer::::new( EchoMcpConnect { context_tx, - notice_tx, runner_started: runner_started.clone(), dropped_tx: Mutex::new(Some(connector_dropped_tx)), }, @@ -793,7 +750,7 @@ async fn v2_resume_mcp_attachment_preserves_request_and_lives_for_connection() - "the MCP runner must be first-polled before session/resume is published" ); - assert_round_trip(1, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(1, &mut round_trip_rx, &mut context_rx).await?; let opened = pending_session.block_task().await?; assert_eq!(opened.session().session_id(), &session_id); @@ -806,7 +763,7 @@ async fn v2_resume_mcp_attachment_preserves_request_and_lives_for_connection() - round_trip_trigger_tx .unbounded_send(()) .map_err(Error::into_internal_error)?; - assert_round_trip(2, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(2, &mut round_trip_rx, &mut context_rx).await?; Ok(()) }) From 8662cbe99701ea75933821dc0a263fce6c281a00 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Thu, 24 Sep 2026 15:24:38 +0200 Subject: [PATCH 2/8] feat(acp): validate stateless MCP transport end to end --- Cargo.lock | 4 +- Cargo.toml | 1 + README.md | 5 + justfile | 9 +- md/SUMMARY.md | 1 + md/mcp-bridge.md | 80 +- md/mcp-over-acp.md | 89 + md/protocol.md | 115 +- .../src/trace.rs | 99 +- .../tests/mcp_over_acp_polyfill.rs | 73 +- .../tests/mcp_over_acp_polyfill_v2.rs | 722 +++---- .../tests/mcp_server_handler_chain_v2.rs | 53 +- .../tests/request_cancellation.rs | 408 +++- .../tests/scoped_mcp_server.rs | 4 +- .../tests/standalone_mcp_server.rs | 2 +- .../tests/test_mcp_connection_context.rs | 25 +- .../tests/test_tool_fn.rs | 2 +- .../tests/trace_client_mcp_server.rs | 31 +- .../tests/trace_mcp_tool_call.rs | 180 +- src/agent-client-protocol-cookbook/src/lib.rs | 7 +- .../CHANGELOG.md | 15 + src/agent-client-protocol-polyfill/Cargo.toml | 7 +- .../src/mcp_over_acp/actor.rs | 78 - .../src/mcp_over_acp/http.rs | 1784 +++++------------ .../src/mcp_over_acp/mod.rs | 1322 +++++------- .../src/mcp_over_acp/protocol.rs | 611 ++---- src/agent-client-protocol-rmcp/Cargo.toml | 5 + src/agent-client-protocol-rmcp/README.md | 14 +- .../examples/stateless_native_mcp.rs | 123 ++ .../tests/stateless_native_mcp.rs | 320 +++ src/agent-client-protocol-test/src/testy.rs | 30 +- src/agent-client-protocol/CHANGELOG.md | 12 + src/agent-client-protocol/src/jsonrpc.rs | 2 +- .../src/mcp_server/active_session.rs | 294 ++- 34 files changed, 3021 insertions(+), 3506 deletions(-) create mode 100644 md/mcp-over-acp.md delete mode 100644 src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs create mode 100644 src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs create mode 100644 src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs diff --git a/Cargo.lock b/Cargo.lock index 19843773..3bce6c5f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -110,11 +110,9 @@ dependencies = [ "agent-client-protocol", "async-stream", "axum", + "base64 0.23.1", "futures", - "futures-concurrency", - "rustc-hash", "serde_json", - "thiserror", "tokio", "tracing", "uuid", diff --git a/Cargo.toml b/Cargo.toml index 4c55ba77..141e7efc 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -45,6 +45,7 @@ tokio-util = { version = "0.7", features = ["compat"] } async-tungstenite = { version = "0.35.0", default-features = false, features = ["tokio-rustls-webpki-roots"] } # Serialization +base64 = "0.23" serde = { version = "1.0", features = ["derive", "rc"] } serde_json = { version = "1", features = ["preserve_order", "raw_value"] } schemars = { version = "1.0", features = ["derive"] } diff --git a/README.md b/README.md index f62a2572..6dafef29 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,11 @@ attaches one while forking. Successful v2 attachments remain active for the connection lifetime, and all three builders expose `on_proxy_session_start` to forward proxied setup without coupling later session events to that response. +The native transport targets MCP 2026-07-28: requests carry their own metadata +and logical IDs, with request-scoped notifications and cancellation rather than +an MCP connection lifecycle. See [Native MCP-over-ACP](./md/mcp-over-acp.md) +for the direct rmcp example and current resource-limit caveats. + **Proxy orchestration** - [`agent-client-protocol-conductor`](./src/agent-client-protocol-conductor/) – Binary and library that manages chains of proxy components. diff --git a/justfile b/justfile index 6dd8c25d..314805e7 100644 --- a/justfile +++ b/justfile @@ -1,9 +1,12 @@ +# Keep file-based snapshots inside this checkout, even in nested worktrees. +export CARGO_WORKSPACE_DIR := justfile_directory() + # Build binaries needed for integration tests prep-tests: cargo build -p agent-client-protocol-conductor --all-features cargo build -p agent-client-protocol-test --bin testy --all-features cargo build -p agent-client-protocol-test --bin mcp-echo-server --example arrow_proxy --all-features -# Run all tests (requires prep-tests first) -test: prep-tests - cargo test --all --workspace --all-features +# Run all tests, or pass a test-name filter / cargo test arguments. +test *args: prep-tests + cargo test --all --workspace --all-features {{args}} diff --git a/md/SUMMARY.md b/md/SUMMARY.md index c51ce8a0..0e619d4a 100644 --- a/md/SUMMARY.md +++ b/md/SUMMARY.md @@ -18,6 +18,7 @@ - [Transport Architecture](./transport-architecture.md) - [HTTP / WebSocket Transport](./http-transport.md) +- [Native MCP-over-ACP](./mcp-over-acp.md) # Conductor (agent-client-protocol-conductor) diff --git a/md/mcp-bridge.md b/md/mcp-bridge.md index 7caddd40..74fa9f48 100644 --- a/md/mcp-bridge.md +++ b/md/mcp-bridge.md @@ -1,14 +1,18 @@ -# MCP-over-ACP Compatibility Bridge +# Stateless MCP-over-ACP HTTP Adapter `agent-client-protocol-polyfill::mcp_over_acp::McpOverAcpPolyfill` adapts the -native ACP MCP transport for a final agent that accepts HTTP MCP -servers. MCP adaptation is explicit and is not built into the conductor. +native ACP MCP transport for a final agent with an MCP 2026-07-28 HTTP client. +MCP adaptation is explicit and is not built into the conductor. There is no +fallback to older MCP revisions. The component-facing side of the bridge always uses the opt-in native protocol: - Servers are declared as `McpServer::Acp` with a `serverId`. -- Connections use `mcp/connect`, `mcp/message`, and `mcp/disconnect`. -- `mcp/disconnect` is a request with a response. +- Each operation uses `mcp/message` with `serverId` and a logical `requestId`. +- The provider sends notifications for that operation; the final ACP response + carries its MCP result or error. +- ACP request cancellation stops only that operation. There is no MCP + initialize/connect/disconnect or session-header exchange. The SDK-local underscore-prefixed method family and HTTP declarations with a special URL scheme have been retired. The polyfill now translates native @@ -93,12 +97,14 @@ polyfill: final agent. 2. Retains the native `serverId` so connections can be routed back to the component that provided the server. -3. Opens the endpoint's native connection by sending `mcp/connect` with that - server ID toward the provider. -4. Relays requests and notifications through `mcp/message`, using the returned - `connectionId` for that active MCP connection. -5. Sends an `mcp/disconnect` request when the local transport closes and removes - the connection from the bridge. +3. Adds a runtime-only bearer credential to the HTTP declaration. The endpoint + requires that credential and checks supplied Origin headers; an ephemeral + port alone is not access control. +4. For each POST, allocates a unique logical MCP request ID and sends + `mcp/message` to the provider. Two HTTP clients may use the same external + JSON-RPC ID without sharing routing or state. +5. Relays notifications and a final result/error for that request. Closing + its HTTP response cancels the corresponding ACP request, not the listener. Enable the polyfill crate's `unstable_session_fork` feature when adapting fork requests. Stable v1 setup includes `session/new`, `session/load`, and @@ -120,27 +126,51 @@ Reference](./protocol.md#native-mcp-over-acp). `McpOverAcpPolyfill::http()` is the default compatibility shape. It replaces the native declaration with an HTTP MCP URL at `http://127.0.0.1:PORT`. The -embedded server accepts MCP POST requests and an SSE GET stream at `/`, retaining -JSON-RPC batch frames and correlating each POST with its response. +embedded server accepts a single JSON-RPC request per POST at `/`, returning +JSON for a terminal-only response or SSE for a request that emits notifications. +GET and DELETE return 405. Batches and client-originated JSON-RPC responses +are rejected; there is no standalone GET event stream or MCP session ID. ```rust,ignore let bridge = McpOverAcpPolyfill::http(); ``` -The listener is bound only on loopback and uses an ephemeral port. It does not -implement resumable SSE event IDs. +Clients must send the bearer header from the declaration, both JSON and SSE +Accept types, and the required MCP protocol-version, method, and applicable +name headers. Mirrored names support MCP's Base64 sentinel encoding. Missing, +duplicate, or mismatched routing headers are rejected. + +The listener is bound only on loopback. Resumable SSE event IDs are not part of +the target MCP revision. Subscription IDs inside +`_meta["io.modelcontextprotocol/subscriptionId"]` are translated back to the +HTTP request's original ID in notifications and graceful completion results; +other metadata, progress tokens, and opaque retry state are not rewritten. ## Lifecycle and Failure Behavior -Each bridge endpoint receives a unique `connectionId` from `mcp/connect`. The -polyfill keeps a connection map until the endpoint's transport task closes, -then removes the entry, sends `mcp/disconnect`, and observes its response. -Request failures use the corresponding request's error path; notifications are -never answered with synthetic errors. +Each POST owns a pending native request, not an MCP session. A terminal result, +error, response-stream close, or overflow removes that request's routing state. +The listening endpoint remains available for later requests. + +The adapter limits each response's queued notifications to 16 messages and +256 KiB of serialized data, with 64 active requests and 32 listening endpoints +per adapter. A separate terminal-response path avoids stranding completion +behind a full queue. Overflow explicitly fails and cancels that operation +without blocking the shared runner or dropping events silently. + +Unknown or late provider notifications are ignored; reverse MCP requests are +not supported. The adapter does not infer ACP session IDs or maintain MCP +initialization state. + +## Remaining scope -A reverse `mcp/message` request for an unknown `connectionId` receives -`Invalid params`. A reverse notification for an unknown connection is ignored, -as required for JSON-RPC notifications. +Tools using `x-mcp-header` annotations are currently unsupported and fail +closed: they are omitted from listings, calls are rejected, and supplied +`Mcp-Param-*` headers are rejected. For a direct tool call the adapter fetches +the tool descriptor internally, including pagination, so the caller does not +need a prior tools/list handshake. That lookup is an explicit per-call cost. -The polyfill does not infer or store ACP session IDs. Association is carried by -the declared `serverId` and the resulting active `connectionId`. +This is not yet full HTTP conformance. Native SDK `Channel` and outgoing +queues also remain unbounded; the HTTP queue limits above do not establish +end-to-end native backpressure. Native admission/payload limits and the +remaining transport work are described in [Native MCP-over-ACP](./mcp-over-acp.md). diff --git a/md/mcp-over-acp.md b/md/mcp-over-acp.md new file mode 100644 index 00000000..190ea7ae --- /dev/null +++ b/md/mcp-over-acp.md @@ -0,0 +1,89 @@ +# Native MCP-over-ACP + +The native transport targets MCP 2026-07-28 only. It lets an ACP client or proxy +provide MCP tools to an agent over the existing ACP connection, without a +conductor, HTTP listener, subprocess, or MCP initialization handshake. + +Enable `unstable_mcp_over_acp` on the core SDK. Draft ACP v2 additionally +requires `unstable_protocol_v2`. The shared-schema revision is currently pinned +to a Git commit for cross-repository validation; replace that pin with the +released schema before publishing the SDK. + +## Providing tools + +Attach an `mcp_server::McpServer` to session setup through the existing builder +APIs. It publishes a `McpServer::Acp` declaration with a provider-generated +`serverId`. + +Each incoming `mcp/message` invokes the backend factory for one operation. +The MCP request context exposes `server_id()` and `request_id()`; standalone +MCP serving has neither. Tool definitions can be shared, but per-request MCP +metadata and capabilities must not be inferred from previous operations. + +The rmcp integration can construct tools through its builder or wrap a supplied +rmcp 3.4 service. The normal rmcp service can process a modern request without +`initialize` when its inner `_meta` declares the modern version and capabilities. + +## Consuming tools + +An ACP agent holds a `ConnectionTo` or its v2 counterpart. It sends +`MessageMcpRequest::new(server_id, request_id, method)` with the inner MCP +parameters, including: + +- `io.modelcontextprotocol/protocolVersion: "2026-07-28"`; +- `io.modelcontextprotocol/clientCapabilities` as an object; +- any request-specific identity, progress token, extension settings, or retry + state required by the MCP operation. + +Choose a fresh logical request ID. It becomes the MCP JSON-RPC ID and remains +unchanged through proxies. The outer ACP request ID is separate and may change +on each hop. + +Register a `MessageMcpNotification` handler before sending requests that may +stream notifications. Route by server and logical request ID. Do not block +the ACP dispatch loop waiting for peer traffic; use a spawned task or the +connection's application future. + +The final response is the MCP result directly, including its `resultType`, or +the original MCP error. For MRTR, process the `input_required` result and send +a fresh request with `inputResponses` and the exact opaque `requestState`. + +Discovery reports only the MCP revision exposed by this binding, even if the +hosted backend also supports older revisions through other transports. + +## Subscriptions and cancellation + +`subscriptions/listen` keeps one request alive. Its acknowledgement and updates +arrive as request-scoped notifications, with the logical request ID in +`io.modelcontextprotocol/subscriptionId`. An unrelated tool call does not share +that subscription's state or lifetime. + +Use `SentRequest::cancel` (or drop an unconsumed request) to cancel the outer +ACP operation. The provider stops that operation's backend work and returns a +result or cancellation error. Removing a provider stops its outstanding work; +no separate `mcp/disconnect` exchange exists. + +## Resource limits and remaining work + +The native provider admits at most 64 concurrent operations per declared +server and checks a 16 MiB serialized payload limit before starting work or +forwarding backend responses/notifications. Rejected work reports an error; +completion and cancellation release the admission slot. + +These are not end-to-end memory bounds. The public SDK `Channel` and outgoing +queues remain unbounded. A bounded native transport path is still required +before stabilization; admission and per-message size checks do not prevent +accumulation behind a slow peer. The [HTTP adapter](./mcp-bridge.md) separately +bounds its own response queues and fails/cancels an overflowing operation. + +## Runnable example + +```sh +cargo run -p agent-client-protocol-rmcp \ + --example stateless_native_mcp \ + --features unstable_mcp_over_acp,unstable_protocol_v2 +``` + +This direct ACP example uses actual rmcp tools without the HTTP polyfill. +See the [protocol reference](./protocol.md#native-mcp-over-acp) for wire details +and the [RFD](https://agentclientprotocol.com/rfds/mcp-over-acp) for the design. diff --git a/md/protocol.md b/md/protocol.md index d30d6747..8b5751cd 100644 --- a/md/protocol.md +++ b/md/protocol.md @@ -11,9 +11,7 @@ unstable and is available only with the `unstable_mcp_over_acp` feature. | --- | --- | --- | | `_proxy/initialize` | request | Initialize a component as a proxy | | `_proxy/successor` | request or notification | Forward one inner ACP message to the next component | -| `mcp/connect` | request | Open a connection to an ACP-provided MCP server | -| `mcp/message` | request or notification | Carry one inner MCP message over ACP | -| `mcp/disconnect` | request | Close an MCP-over-ACP connection | +| `mcp/message` | agent request or provider notification | Invoke an MCP operation or carry a notification for that operation | There are no separate request and notification method names for successor or MCP message forwarding. The presence of an outer JSON-RPC `id` distinguishes a @@ -61,10 +59,11 @@ inner message. ## Native MCP-over-ACP -Enable `unstable_mcp_over_acp` to use the draft native transport. A component -providing an MCP server adds `McpServer::Acp` to session setup requests -(`session/new`, `session/load`, `session/resume`, and the opt-in `session/fork`). -Its wire shape contains a human-readable name and an opaque server identifier: +Enable `unstable_mcp_over_acp` to use the draft native transport targeting MCP +2026-07-28 only. ACP initialization is unchanged; there is no MCP initialization +or connect/disconnect lifecycle. A provider adds `McpServer::Acp` to session +setup requests (`session/new`, `session/resume`, v1 `session/load`, and the +opt-in `session/fork`): ```json { @@ -74,48 +73,22 @@ Its wire shape contains a human-readable name and an opaque server identifier: } ``` -`serverId` identifies the declared server and is used to route `mcp/connect` +`serverId` identifies the declared server and is used to route `mcp/message` back to the component that provided it. A provider must not reuse one server ID for multiple visible servers on the same ACP connection. The high-level `agent_client_protocol::mcp_server::McpServer` APIs create this declaration automatically. -An agent that consumes this transport advertises -`agentCapabilities.mcpCapabilities.acp`. If the final agent supports HTTP but -not ACP-transport MCP servers, place the [MCP-over-ACP compatibility -bridge](./mcp-bridge.md) immediately before it. - -### `mcp/connect` - -The MCP client opens a connection to the declared server ID: - -```json -{ - "jsonrpc": "2.0", - "id": 20, - "method": "mcp/connect", - "params": { "serverId": "mcp-server:01" } -} -``` - -The provider creates one active MCP connection and returns a distinct -connection ID: - -```json -{ - "jsonrpc": "2.0", - "id": 20, - "result": { "connectionId": "mcp-connection:01" } -} -``` - -The server ID selects what to connect to; the connection ID selects that -particular running connection. All subsequent messages use the connection ID. +An agent advertises `agentCapabilities.mcpCapabilities.acp: true` in v1 or +`capabilities.session.mcp.acp: {}` in draft v2. An optional +[HTTP adapter](./mcp-bridge.md) is only for agents with a modern MCP HTTP client. +Advertising HTTP support alone does not establish MCP revision compatibility. ### `mcp/message` -`mcp/message` carries one inner MCP method and its named parameters. The method -is bidirectional because MCP clients and servers can both issue requests: +An agent sends one request addressed to the server, with a fresh logical MCP +request ID. This ID remains unchanged through proxies even if the outer ACP +JSON-RPC ID is renumbered: ```json { @@ -123,46 +96,68 @@ is bidirectional because MCP clients and servers can both issue requests: "id": 21, "method": "mcp/message", "params": { - "connectionId": "mcp-connection:01", + "serverId": "mcp-server:01", + "requestId": "mcp-request:01", "method": "tools/call", "params": { "name": "example", - "arguments": {} + "arguments": {}, + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + "progressToken": "caller-supplied-token" + } } } } ``` -Use an outer request for an inner MCP request and an outer notification for an -inner MCP notification. The outer response carries the inner MCP result or -error. +The outer response carries the inner MCP result (including `resultType`) or +error directly. MRTR `input_required` is a result, not a reverse RPC; retry the +original operation with fresh metadata/IDs and unchanged opaque state. -### `mcp/disconnect` +For `server/discover`, supported versions are restricted to the revision +exposed by this binding; a backend must actually support that revision. -Disconnect is a request so the caller knows that the provider has released the -active connection: +A provider may send notifications belonging to that operation: ```json { "jsonrpc": "2.0", - "id": 22, - "method": "mcp/disconnect", - "params": { "connectionId": "mcp-connection:01" } + "method": "mcp/message", + "params": { + "serverId": "mcp-server:01", + "requestId": "mcp-request:01", + "method": "notifications/progress", + "params": { "progressToken": "caller-supplied-token", "progress": 1 } + } } ``` -A successful disconnect returns an empty result: +Progress requires a corresponding token in the original request's inner MCP +metadata. Subscription notifications carry the listen request's logical +`requestId` in `io.modelcontextprotocol/subscriptionId`; acknowledgement comes +first. Notifications stop when their operation completes. -```json -{ - "jsonrpc": "2.0", - "id": 22, - "result": {} -} -``` +Both envelope types require non-null `serverId`, `requestId`, and `method` +strings. Inner `params` accepts an object or `null`; omission and `null` both +mean no parameters. A valid modern request still needs its required +`params._meta`. Optional outer ACP `_meta` is distinct from inner MCP metadata. + +### Cancellation and lifetime + +Use [`$/cancel_request`](./request-cancellation.md) with the outer ACP request +ID. Normal proxy forwarding maps this cancellation hop by hop. It never +rewrites the logical MCP ID. + +Each operation owns its backend work. A result, error, cancellation, or +provider removal ends that operation; sibling requests and subscriptions stay +independent. There is no MCP connection ID to release. `server/discover` is an +ordinary optional request, not a prerequisite for tool calls. ## Related Documentation +- [Native MCP-over-ACP](./mcp-over-acp.md) - [Conductor Design](./conductor.md) - [MCP Bridge](./mcp-bridge.md) - [Original P/ACP Design Proposal](./proxying-acp.md) (historical) diff --git a/src/agent-client-protocol-conductor/src/trace.rs b/src/agent-client-protocol-conductor/src/trace.rs index 62763290..3921e954 100644 --- a/src/agent-client-protocol-conductor/src/trace.rs +++ b/src/agent-client-protocol-conductor/src/trace.rs @@ -232,8 +232,9 @@ impl TraceWriter { id: serde_json::Value, method: String, session: Option, - params: serde_json::Value, + mut params: serde_json::Value, ) { + redact_http_credentials(&mut params); self.request_details.insert( id.clone(), RequestDetails { @@ -262,8 +263,9 @@ impl TraceWriter { to: ComponentIndex, id: serde_json::Value, is_error: bool, - payload: serde_json::Value, + mut payload: serde_json::Value, ) { + redact_http_credentials(&mut payload); self.write_event(&TraceEvent::Response(ResponseEvent { ts: self.elapsed(), from: format!("{from:?}"), @@ -282,8 +284,9 @@ impl TraceWriter { to: ComponentIndex, method: impl Into, session: Option, - params: serde_json::Value, + mut params: serde_json::Value, ) { + redact_http_credentials(&mut params); self.write_event(&TraceEvent::Notification(NotificationEvent { ts: self.elapsed(), protocol, @@ -526,6 +529,58 @@ fn params_from_transport(params: Option) -> serde_json::Value params.map_or(serde_json::Value::Null, RawJsonRpcParams::into_value) } +/// Do not persist HTTP credentials from MCP declarations or other traced payloads. +/// Only the trace's copy is modified; transport messages retain their headers. +fn redact_http_credentials(value: &mut serde_json::Value) { + fn is_credential(name: &str) -> bool { + [ + "authorization", + "proxy-authorization", + "cookie", + "set-cookie", + "x-api-key", + ] + .iter() + .any(|candidate| name.eq_ignore_ascii_case(candidate)) + } + + match value { + serde_json::Value::Object(object) => { + match object.get_mut("headers") { + Some(serde_json::Value::Array(headers)) => { + for header in headers { + if header + .get("name") + .and_then(serde_json::Value::as_str) + .is_some_and(is_credential) + && let Some(value) = header.get_mut("value") + { + *value = serde_json::Value::String("[REDACTED]".to_owned()); + } + } + } + Some(serde_json::Value::Object(headers)) => { + for (name, value) in headers { + if is_credential(name) { + *value = serde_json::Value::String("[REDACTED]".to_owned()); + } + } + } + _ => {} + } + for value in object.values_mut() { + redact_http_credentials(value); + } + } + serde_json::Value::Array(values) => { + for value in values { + redact_http_credentials(value); + } + } + _ => {} + } +} + /// A message observed going over a channel connected to `left` and `right`. /// This could be a successor message, a mcp-over-acp message, etc. #[derive(Debug)] @@ -651,16 +706,46 @@ mod tests { use agent_client_protocol::RawJsonRpcMessage; use serde_json::json; - use super::{MessageInfo, Protocol}; + use super::{MessageInfo, Protocol, redact_http_credentials}; + + #[test] + fn trace_credentials_are_redacted_in_nested_header_shapes() { + let original = json!({ + "params": { + "mcpServers": [{ + "type": "http", + "headers": [ + {"name": "Authorization", "value": "Bearer test-token"}, + {"name": "X-Trace-Id", "value": "keep"}, + {"name": "cOoKiE", "value": "test-cookie"} + ] + }] + }, + "other": {"headers": {"X-Api-Key": "test-key", "Accept": "application/json"}} + }); + let mut traced = original.clone(); + redact_http_credentials(&mut traced); + let headers = &traced["params"]["mcpServers"][0]["headers"]; + assert_eq!(headers[0]["value"], "[REDACTED]"); + assert_eq!(headers[1]["value"], "keep"); + assert_eq!(headers[2]["value"], "[REDACTED]"); + assert_eq!(traced["other"]["headers"]["X-Api-Key"], "[REDACTED]"); + assert_eq!(traced["other"]["headers"]["Accept"], "application/json"); + assert_eq!( + original["params"]["mcpServers"][0]["headers"][0]["value"], + "Bearer test-token" + ); + } #[test] - fn tolerant_mcp_notification_params_are_traced_as_mcp() { + fn nullable_mcp_notification_params_are_traced_as_mcp() { let RawJsonRpcMessage::Notification(notification) = RawJsonRpcMessage::notification( "mcp/message".into(), json!({ - "connectionId": "connection-1", + "serverId": "server-1", + "requestId": "request-1", "method": "notifications/progress", - "params": ["invalid named params"] + "params": null }), ) .expect("notification is valid JSON-RPC") else { diff --git a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs index e2adc450..41927ada 100644 --- a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs +++ b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs @@ -6,15 +6,15 @@ use std::sync::{Arc, Mutex}; use agent_client_protocol::schema::ProtocolVersion; use agent_client_protocol::schema::v1::{ - AgentCapabilities, ConnectMcpRequest, ConnectMcpResponse, InitializeRequest, - InitializeResponse, LoadSessionRequest, LoadSessionResponse, McpCapabilities, McpServer, - McpServerAcp, NewSessionRequest, NewSessionResponse, ResumeSessionRequest, + AgentCapabilities, InitializeRequest, InitializeResponse, LoadSessionRequest, + LoadSessionResponse, McpCapabilities, McpServer, McpServerAcp, MessageMcpRequest, + MessageMcpResponse, NewSessionRequest, NewSessionResponse, ResumeSessionRequest, ResumeSessionResponse, SessionCapabilities, SessionResumeCapabilities, }; use agent_client_protocol::{Agent, Client, Conductor, ConnectTo, Proxy}; use agent_client_protocol_conductor::{ConductorImpl, ProxiesAndAgent}; use agent_client_protocol_polyfill::mcp_over_acp::McpOverAcpPolyfill; -use tokio::io::duplex; +use tokio::io::{AsyncReadExt, AsyncWriteExt, duplex}; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; const SERVER_NAME: &str = "shared-server"; @@ -56,7 +56,7 @@ struct RecordingAgent { } struct NativeMcpProvider { - connect_count: Arc, + request_count: Arc, } impl ConnectTo for NativeMcpProvider { @@ -69,10 +69,12 @@ impl ConnectTo for NativeMcpProvider { .name("native-mcp-provider") .on_receive_request_from( Agent, - async move |request: ConnectMcpRequest, responder, _cx| { + async move |request: MessageMcpRequest, responder, _cx| { assert_eq!(request.server_id.to_string(), SERVER_ID); - self.connect_count.fetch_add(1, Ordering::SeqCst); - responder.respond(ConnectMcpResponse::new("test-connection-id")) + self.request_count.fetch_add(1, Ordering::SeqCst); + responder.respond(serde_json::from_value::( + serde_json::json!({"tools": []}), + )?) }, agent_client_protocol::on_receive_request!(), ) @@ -144,6 +146,23 @@ fn native_server() -> McpServer { McpServer::Acp(McpServerAcp::new(SERVER_NAME, SERVER_ID).meta(meta)) } +async fn http_post(url: &str, bearer: &str, id: i64) -> serde_json::Value { + let address = url.strip_prefix("http://").unwrap(); + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + let body = serde_json::json!({"jsonrpc":"2.0","id":id,"method":"tools/list", + "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28"}}}) + .to_string(); + let request = format!( + "POST / HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: tools/list\r\nContent-Length: {}\r\n\r\n{body}", + body.len() + ); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut response = String::new(); + stream.read_to_string(&mut response).await.unwrap(); + assert!(response.starts_with("HTTP/1.1 200"), "{response}"); + serde_json::from_str(response.split("\r\n\r\n").nth(1).unwrap()).unwrap() +} + async fn recv( response: agent_client_protocol::SentRequest, ) -> Result { @@ -158,7 +177,7 @@ async fn recv( async fn run_with_polyfill( agent: RecordingAgent, - provider_connect_count: Arc, + provider_request_count: Arc, editor_task: impl AsyncFnOnce( agent_client_protocol::ConnectionTo, ) -> Result<(), agent_client_protocol::Error>, @@ -184,7 +203,7 @@ async fn run_with_polyfill( "polyfill-test-conductor".to_string(), ProxiesAndAgent::new(agent) .proxy(NativeMcpProvider { - connect_count: provider_connect_count, + request_count: provider_request_count, }) .proxy(McpOverAcpPolyfill::http()), ) @@ -206,9 +225,9 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ capabilities: agent_capabilities(McpCapabilities::new().http(true)), observed: observed.clone(), }; - let connect_count = Arc::new(AtomicUsize::new(0)); + let request_count = Arc::new(AtomicUsize::new(0)); - run_with_polyfill(agent, connect_count.clone(), async |connection| { + run_with_polyfill(agent, request_count.clone(), async |connection| { let initialize = recv(connection.send_request(InitializeRequest::new(ProtocolVersion::V1))).await?; assert!(initialize.agent_capabilities.mcp_capabilities.http); @@ -235,6 +254,20 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ )) .await?; + let (url, bearer) = { + let setup = observed.setup.lock().unwrap(); + let McpServer::Http(server) = &setup[0].mcp_servers[0] else { + panic!("expected HTTP declaration") + }; + (server.url.clone(), server.headers[0].value.clone()) + }; + let (first, second) = + tokio::join!(http_post(&url, &bearer, 1), http_post(&url, &bearer, 1),); + assert_eq!( + first, + serde_json::json!({"jsonrpc":"2.0","id":1,"result":{"tools":[]}}) + ); + assert_eq!(second, first); Ok(()) }) .await?; @@ -244,9 +277,9 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ .lock() .expect("setup request mutex should not be poisoned"); assert_eq!( - connect_count.load(Ordering::SeqCst), - 1, - "one reused listener should create one native MCP connection" + request_count.load(Ordering::SeqCst), + 2, + "each HTTP POST creates exactly one native MCP request, without a connect handshake" ); assert_eq!(setup.len(), 3); assert_eq!(setup[0].method, SetupMethod::New); @@ -267,7 +300,9 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ }; assert_eq!(server.name, SERVER_NAME); assert_eq!(server.meta.as_ref(), Some(&expected_meta)); - assert!(server.headers.is_empty()); + assert_eq!(server.headers.len(), 1); + assert_eq!(server.headers[0].name, "Authorization"); + assert!(server.headers[0].value.starts_with("Bearer ")); assert!(server.url.starts_with("http://127.0.0.1:")); if let Some(endpoint) = &endpoint { assert_eq!( @@ -292,9 +327,9 @@ async fn native_downstream_keeps_capability_and_declaration_unchanged() }; let declaration = native_server(); let expected = declaration.clone(); - let connect_count = Arc::new(AtomicUsize::new(0)); + let request_count = Arc::new(AtomicUsize::new(0)); - run_with_polyfill(agent, connect_count.clone(), async move |connection| { + run_with_polyfill(agent, request_count.clone(), async move |connection| { let initialize = recv(connection.send_request(InitializeRequest::new(ProtocolVersion::V1))).await?; assert!(!initialize.agent_capabilities.mcp_capabilities.http); @@ -315,7 +350,7 @@ async fn native_downstream_keeps_capability_and_declaration_unchanged() assert_eq!(setup.len(), 1); assert_eq!(setup[0].mcp_servers, vec![expected]); assert_eq!( - connect_count.load(Ordering::SeqCst), + request_count.load(Ordering::SeqCst), 0, "a native-capable downstream should not be routed through the HTTP adapter" ); diff --git a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs index 284ee0ce..a7d10dc9 100644 --- a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs +++ b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs @@ -1,9 +1,8 @@ #![cfg(feature = "unstable_protocol_v2")] -//! V2 integration coverage for the public MCP-over-ACP compatibility proxy. +//! End-to-end v2 coverage for the request-scoped MCP HTTP adapter. use std::{ - collections::BTreeMap, path::PathBuf, sync::{ Arc, Mutex, @@ -17,71 +16,32 @@ use agent_client_protocol::{ }; use agent_client_protocol_conductor::{ConductorImpl, ProxiesAndAgent}; use agent_client_protocol_polyfill::mcp_over_acp::McpOverAcpPolyfill; -use rmcp::{ - ServiceExt as _, - transport::{ - StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig, - }, -}; -use tokio::io::duplex; +use tokio::io::{AsyncReadExt, AsyncWriteExt, duplex}; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; -const SERVER_NAME: &str = "shared-v2-server"; -const SERVER_ID: &str = "shared-v2-server-id"; +const SERVER_ID: &str = "v2-server-id"; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -enum SetupMethod { - New, - Resume, -} - -#[derive(Debug)] -struct SetupRequest { - method: SetupMethod, - mcp_servers: Vec, -} - -#[derive(Default)] -struct ObservedRequests { - setup: Mutex>, -} - -impl ObservedRequests { - fn record(&self, method: SetupMethod, mcp_servers: Vec) { - self.setup - .lock() - .expect("setup request mutex should not be poisoned") - .push(SetupRequest { - method, - mcp_servers, - }); - } -} - -struct RecordingAgent { +struct TestAgent { capabilities: v2::AgentCapabilities, - observed: Arc, + observed: Arc>>, } -impl ConnectTo for RecordingAgent { +impl ConnectTo for TestAgent { async fn connect_to( self, client: impl ConnectTo, ) -> Result<(), agent_client_protocol::Error> { let capabilities = self.capabilities; - let new_observed = Arc::clone(&self.observed); - let resume_observed = self.observed; - + let observed = self.observed; Agent .v2() - .name("recording-v2-agent") + .name("v2-http-test-agent") .on_receive_request( async move |request: v2::InitializeRequest, responder, _cx| { - assert_eq!(request.protocol_version, ProtocolVersion::V2); responder.respond( v2::InitializeResponse::new( request.protocol_version, - implementation("recording-v2-agent"), + v2::Implementation::new("test", "1.0.0"), ) .capabilities(capabilities.clone()), ) @@ -90,15 +50,8 @@ impl ConnectTo for RecordingAgent { ) .on_receive_request( async move |request: v2::NewSessionRequest, responder, _cx| { - new_observed.record(SetupMethod::New, request.mcp_servers); - responder.respond(v2::NewSessionResponse::new("v2-session-id")) - }, - agent_client_protocol::on_receive_request!(), - ) - .on_receive_request( - async move |request: v2::ResumeSessionRequest, responder, _cx| { - resume_observed.record(SetupMethod::Resume, request.mcp_servers); - responder.respond(v2::ResumeSessionResponse::new()) + *observed.lock().unwrap() = request.mcp_servers; + responder.respond(v2::NewSessionResponse::new("session")) }, agent_client_protocol::on_receive_request!(), ) @@ -107,85 +60,72 @@ impl ConnectTo for RecordingAgent { } } -struct NativeMcpProvider { - connect_count: Arc, - request_methods: Arc>>, - notification_methods: Arc>>, - disconnect_count: Arc, -} +struct TestProvider(Arc>>, Arc, Arc); -impl ConnectTo for NativeMcpProvider { +impl ConnectTo for TestProvider { async fn connect_to( self, client: impl ConnectTo, ) -> Result<(), agent_client_protocol::Error> { - let request_methods = Arc::clone(&self.request_methods); - let notification_methods = Arc::clone(&self.notification_methods); - let disconnect_count = Arc::clone(&self.disconnect_count); - Proxy .v2() - .name("native-v2-mcp-provider") + .name("v2-mcp-provider") .on_receive_request_from( Agent, - async move |request: v2::ConnectMcpRequest, responder, _cx| { + async move |request: v2::MessageMcpRequest, responder, cx| { assert_eq!(request.server_id.to_string(), SERVER_ID); - self.connect_count.fetch_add(1, Ordering::SeqCst); - responder.respond(v2::ConnectMcpResponse::new("v2-test-connection-id")) - }, - agent_client_protocol::on_receive_request!(), - ) - .on_receive_request_from( - Agent, - async move |request: v2::MessageMcpRequest, responder, _cx| { - request_methods - .lock() - .expect("request method mutex should not be poisoned") - .push(request.method.clone()); - match request.method.as_str() { - "initialize" => { - let protocol_version = request - .params - .as_ref() - .and_then(|params| params.get("protocolVersion")) - .cloned() - .unwrap_or_else(|| serde_json::json!("2025-06-18")); - responder.respond(serde_json::from_value(serde_json::json!({ - "protocolVersion": protocol_version, - "capabilities": { - "tools": {} - }, - "serverInfo": { - "name": "v2-polyfill-test-mcp-server", - "version": env!("CARGO_PKG_VERSION") - } - }))?) - } - "tools/list" => responder - .respond(serde_json::from_value(serde_json::json!({ "tools": [] }))?), - method => responder.respond_with_error( - agent_client_protocol::Error::method_not_found().data(method), - ), + self.0.lock().unwrap().push(request.request_id.to_string()); + self.1.fetch_add(1, Ordering::SeqCst); + if request.method == "subscriptions/listen" + || request.method == "subscriptions/flood" + { + let params = if request.method == "subscriptions/flood" { + serde_json::Map::from_iter([( + "payload".into(), + serde_json::json!("x".repeat(300 * 1024)), + )]) + } else { + serde_json::Map::from_iter([( + "_meta".into(), + serde_json::json!({"io.modelcontextprotocol/subscriptionId": + request.request_id.to_string()}), + )]) + }; + cx.send_notification_to( + Agent, + v2::MessageMcpNotification::new( + SERVER_ID, + request.request_id, + "notifications/subscriptions/acknowledged", + ) + .params(params), + )?; + let cancelled = responder.cancellation(); + let count = self.2.clone(); + cx.spawn(async move { + cancelled.cancelled().await; + count.fetch_add(1, Ordering::SeqCst); + responder.respond_with_error( + agent_client_protocol::Error::request_cancelled(), + ) + })?; + return Ok(()); } - }, - agent_client_protocol::on_receive_request!(), - ) - .on_receive_notification_from( - Agent, - async move |notification: v2::MessageMcpNotification, _cx| { - notification_methods - .lock() - .expect("notification method mutex should not be poisoned") - .push(notification.method); - Ok(()) - }, - agent_client_protocol::on_receive_notification!(), - ) - .on_receive_request_from( - Agent, - async move |_request: v2::DisconnectMcpRequest, responder, _cx| { - disconnect_count.fetch_add(1, Ordering::SeqCst); - responder.respond(v2::DisconnectMcpResponse::new()) + let result = match request.method.as_str() { + "tools/list" => serde_json::json!({"tools":[ + {"name":"ping","inputSchema":{"type":"object","properties":{}}}, + {"name":"restricted","inputSchema":{"type":"object","properties":{ + "region":{"type":"string","x-mcp-header":"Region"} + }}} + ]}), + "tools/call" => serde_json::json!({"content":[]}), + _ => { + return responder.respond_with_error( + agent_client_protocol::Error::method_not_found(), + ); + } + }; + responder.respond(serde_json::from_value::(result)?) }, agent_client_protocol::on_receive_request!(), ) @@ -194,78 +134,30 @@ impl ConnectTo for NativeMcpProvider { } } -fn implementation(name: &str) -> v2::Implementation { - v2::Implementation::new(name, env!("CARGO_PKG_VERSION")) -} - -fn agent_capabilities(mcp: v2::McpCapabilities) -> v2::AgentCapabilities { - v2::AgentCapabilities::new().session(v2::SessionCapabilities::new().mcp(mcp)) -} - -fn initialize_request() -> v2::InitializeRequest { - v2::InitializeRequest::new( - ProtocolVersion::V2, - implementation("v2-polyfill-test-client"), - ) -} - -fn server_meta() -> v2::Meta { - let mut meta = v2::Meta::new(); - meta.insert( - "source".to_owned(), - serde_json::Value::String("v2-integration-test".to_owned()), - ); - meta -} - -fn native_server() -> v2::McpServer { - v2::McpServer::Acp(v2::McpServerAcp::new(SERVER_NAME, SERVER_ID).meta(server_meta())) -} - -fn future_server() -> v2::McpServer { - v2::McpServer::Other(v2::OtherMcpServer::new( - "_future_transport", - BTreeMap::from([ - ("name".to_owned(), serde_json::json!("future-v2-server")), - ( - "configuration".to_owned(), - serde_json::json!({ "preserve": true }), - ), - ]), - )) -} - -fn test_servers() -> Vec { - vec![native_server(), future_server()] -} - -async fn run_with_polyfill( - agent: RecordingAgent, - provider_connect_count: Arc, - provider_request_methods: Arc>>, - provider_notification_methods: Arc>>, - provider_disconnect_count: Arc, - editor_task: impl AsyncFnOnce(V2ConnectionTo) -> Result<(), agent_client_protocol::Error>, +async fn run( + capabilities: v2::AgentCapabilities, + observed: Arc>>, + ids: Arc>>, + count: Arc, + cancelled: Arc, + editor: impl AsyncFnOnce(V2ConnectionTo) -> Result<(), agent_client_protocol::Error>, ) -> Result<(), agent_client_protocol::Error> { let (editor_out, conductor_in) = duplex(4096); let (conductor_out, editor_in) = duplex(4096); let transport = agent_client_protocol::ByteStreams::new(editor_out.compat_write(), editor_in.compat()); - Client .v2() - .name("v2-polyfill-test-client") + .name("v2-mcp-test-client") .with_spawned(|_cx| async move { ConductorImpl::new_agent( - "v2-polyfill-test-conductor", - ProxiesAndAgent::new(agent) - .proxy(NativeMcpProvider { - connect_count: provider_connect_count, - request_methods: provider_request_methods, - notification_methods: provider_notification_methods, - disconnect_count: provider_disconnect_count, - }) - .proxy(McpOverAcpPolyfill::http()), + "v2-mcp-test-conductor", + ProxiesAndAgent::new(TestAgent { + capabilities, + observed, + }) + .proxy(TestProvider(ids, count, cancelled)) + .proxy(McpOverAcpPolyfill::http()), ) .run(agent_client_protocol::ByteStreams::new( conductor_out.compat_write(), @@ -273,188 +165,163 @@ async fn run_with_polyfill( )) .await }) - .connect_with(transport, editor_task) + .connect_with(transport, editor) .await } -fn negotiated_mcp_capabilities(response: &v2::InitializeResponse) -> &v2::McpCapabilities { - response - .capabilities - .session - .as_ref() - .expect("the test agent should advertise session support") - .mcp - .as_ref() - .expect("the test agent should advertise MCP support") +fn native_server() -> v2::McpServer { + let mut meta = v2::Meta::new(); + meta.insert("preserve".into(), serde_json::json!(true)); + v2::McpServer::Acp(v2::McpServerAcp::new("native", SERVER_ID).meta(meta)) +} + +fn initialize() -> v2::InitializeRequest { + v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("test", "1.0.0"), + ) } -#[tokio::test] -async fn http_downstream_adapts_v2_capabilities_and_only_transforms_native_servers() --> Result<(), agent_client_protocol::Error> { - let observed = Arc::new(ObservedRequests::default()); - let agent = RecordingAgent { - capabilities: agent_capabilities( - v2::McpCapabilities::new().http(v2::McpHttpCapabilities::new()), - ), - observed: Arc::clone(&observed), +async fn post(url: &str, bearer: &str, method: &str, tool: &str) -> serde_json::Value { + let address = url.strip_prefix("http://").unwrap(); + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + let mut params = serde_json::json!({ + "_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28"} + }); + if method == "tools/call" { + params["name"] = serde_json::json!(tool); + params["arguments"] = serde_json::json!({}); + } + let body = serde_json::json!({"jsonrpc":"2.0","id":"same","method":method, + "params":params}) + .to_string(); + let name = if method == "tools/call" { + format!("Mcp-Name: {tool}\r\n") + } else { + String::new() }; - let connect_count = Arc::new(AtomicUsize::new(0)); - let request_methods = Arc::new(Mutex::new(Vec::new())); - let notification_methods = Arc::new(Mutex::new(Vec::new())); - let disconnect_count = Arc::new(AtomicUsize::new(0)); + let request = format!( + "POST / HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: {method}\r\n{name}Content-Length: {}\r\n\r\n{body}", + body.len() + ); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut response = String::new(); + stream.read_to_string(&mut response).await.unwrap(); + assert!(response.starts_with("HTTP/1.1 200"), "{response}"); + serde_json::from_str(response.split("\r\n\r\n").nth(1).unwrap()).unwrap() +} - run_with_polyfill( - agent, - Arc::clone(&connect_count), - Arc::clone(&request_methods), - Arc::clone(¬ification_methods), - Arc::clone(&disconnect_count), +#[tokio::test] +async fn modern_http_v2_requests_are_stateless_and_isolated() +-> Result<(), agent_client_protocol::Error> { + let observed = Arc::new(Mutex::new(Vec::new())); + let ids = Arc::new(Mutex::new(Vec::new())); + let count = Arc::new(AtomicUsize::new(0)); + let caps = v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new() + .mcp(v2::McpCapabilities::new().http(v2::McpHttpCapabilities::new())), + ); + run( + caps, + observed.clone(), + ids.clone(), + count.clone(), + Arc::default(), async |connection| { - let initialize = connection - .send_request(initialize_request()) - .block_task() - .await?; - let mcp = negotiated_mcp_capabilities(&initialize); - assert!(mcp.http.is_some()); + let initialized = connection.send_request(initialize()).block_task().await?; assert!( - mcp.acp.is_some(), - "the HTTP adapter should advertise v2 native MCP support upstream" + initialized + .capabilities + .session + .unwrap() + .mcp + .unwrap() + .acp + .is_some() ); - - let cwd = PathBuf::from("/tmp"); - let session = connection - .send_request(v2::NewSessionRequest::new(cwd.clone()).mcp_servers(test_servers())) - .block_task() - .await?; connection .send_request( - v2::ResumeSessionRequest::new(session.session_id, cwd) - .mcp_servers(test_servers()), + v2::NewSessionRequest::new(PathBuf::from("/tmp")) + .mcp_servers(vec![native_server()]), ) .block_task() .await?; - - let endpoint = { - let setup = observed - .setup - .lock() - .expect("setup request mutex should not be poisoned"); - let v2::McpServer::Http(server) = &setup[0].mcp_servers[0] else { - panic!("expected the native declaration to be adapted to HTTP") + let (url, bearer) = { + let observed = observed.lock().unwrap(); + let v2::McpServer::Http(server) = &observed[0] else { + panic!("expected HTTP endpoint") }; - server.url.clone() + assert_eq!( + server.meta.as_ref().unwrap().get("preserve"), + Some(&serde_json::json!(true)) + ); + assert_eq!(server.headers[0].name, "Authorization"); + (server.url.clone(), server.headers[0].value.clone()) }; - let mcp_client = () - .serve(StreamableHttpClientTransport::from_config( - StreamableHttpClientTransportConfig::with_uri(endpoint), - )) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - let tools = mcp_client - .list_tools(None) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - assert!(tools.tools.is_empty()); - mcp_client - .cancel() - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; + // No prior client tools/list: the adapter looks up the descriptor + // internally, rejecting annotated tools instead of skipping mirrors. + let direct = post(&url, &bearer, "tools/call", "ping").await; + assert_eq!( + direct, + serde_json::json!({"jsonrpc":"2.0","id":"same","result":{"content":[]}}) + ); + let annotated = post(&url, &bearer, "tools/call", "restricted").await; + assert_eq!(annotated["error"]["code"], -32602); + let (a, b) = tokio::join!( + post(&url, &bearer, "tools/list", ""), + post(&url, &bearer, "tools/list", "") + ); + assert_eq!( + a, + serde_json::json!({"jsonrpc":"2.0","id":"same","result":{"tools":[ + {"name":"ping","inputSchema":{"type":"object","properties":{}}} + ]}}) + ); + assert_eq!(a, b); Ok(()) }, ) .await?; - - let setup = observed - .setup - .lock() - .expect("setup request mutex should not be poisoned"); - assert_eq!( - connect_count.load(Ordering::SeqCst), - 1, - "one reused listener should create one v2 native MCP connection" - ); - assert_eq!( - *request_methods - .lock() - .expect("request method mutex should not be poisoned"), - ["initialize", "tools/list"] - ); - assert_eq!( - *notification_methods - .lock() - .expect("notification method mutex should not be poisoned"), - ["notifications/initialized"] + assert_eq!(count.load(Ordering::SeqCst), 5); + let ids = ids.lock().unwrap(); + assert_eq!(ids.len(), 5); + assert_ne!( + ids[0], ids[1], + "external JSON-RPC IDs must not collide at the ACP hop" ); - assert_eq!(setup.len(), 2); - assert_eq!(setup[0].method, SetupMethod::New); - assert_eq!(setup[1].method, SetupMethod::Resume); - - let expected_future_server = future_server(); - let expected_meta = server_meta(); - let mut endpoint = None; - for request in setup.iter() { - assert_eq!(request.mcp_servers.len(), 2); - let v2::McpServer::Http(server) = &request.mcp_servers[0] else { - panic!( - "expected the ACP declaration to become HTTP for {:?}, got {:?}", - request.method, request.mcp_servers - ); - }; - assert_eq!(server.name, SERVER_NAME); - assert_eq!(server.meta.as_ref(), Some(&expected_meta)); - assert!(server.headers.is_empty()); - assert!(server.url.starts_with("http://127.0.0.1:")); - assert_eq!( - request.mcp_servers[1], expected_future_server, - "the polyfill must preserve custom v2 MCP transports" - ); - if let Some(endpoint) = &endpoint { - assert_eq!( - &server.url, endpoint, - "the same ACP server ID should reuse one listener" - ); - } else { - endpoint = Some(server.url.clone()); - } - } - Ok(()) } #[tokio::test] -async fn native_v2_downstream_keeps_capability_and_declarations_unchanged() +async fn native_v2_declarations_pass_through_without_http() -> Result<(), agent_client_protocol::Error> { - let observed = Arc::new(ObservedRequests::default()); - let agent = RecordingAgent { - capabilities: agent_capabilities( - v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), - ), - observed: Arc::clone(&observed), - }; - let expected = test_servers(); - let connect_count = Arc::new(AtomicUsize::new(0)); - let request_methods = Arc::new(Mutex::new(Vec::new())); - let notification_methods = Arc::new(Mutex::new(Vec::new())); - let disconnect_count = Arc::new(AtomicUsize::new(0)); - - run_with_polyfill( - agent, - Arc::clone(&connect_count), - Arc::clone(&request_methods), - Arc::clone(¬ification_methods), - Arc::clone(&disconnect_count), - async move |connection| { - let initialize = connection - .send_request(initialize_request()) - .block_task() - .await?; - let mcp = negotiated_mcp_capabilities(&initialize); - assert!(mcp.http.is_none()); - assert!(mcp.acp.is_some()); - + let observed = Arc::new(Mutex::new(Vec::new())); + let caps = v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new() + .mcp(v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new())), + ); + run( + caps, + observed.clone(), + Arc::default(), + Arc::default(), + Arc::default(), + async |connection| { + let initialized = connection.send_request(initialize()).block_task().await?; + assert!( + initialized + .capabilities + .session + .unwrap() + .mcp + .unwrap() + .acp + .is_some() + ); connection .send_request( - v2::NewSessionRequest::new(PathBuf::from("/tmp")).mcp_servers(expected.clone()), + v2::NewSessionRequest::new(PathBuf::from("/tmp")) + .mcp_servers(vec![native_server()]), ) .block_task() .await?; @@ -462,105 +329,90 @@ async fn native_v2_downstream_keeps_capability_and_declarations_unchanged() }, ) .await?; - - let setup = observed - .setup - .lock() - .expect("setup request mutex should not be poisoned"); - assert_eq!(setup.len(), 1); - assert_eq!(setup[0].mcp_servers, test_servers()); - assert_eq!( - connect_count.load(Ordering::SeqCst), - 0, - "a native-capable v2 downstream should bypass the HTTP adapter" - ); - assert!( - request_methods - .lock() - .expect("request method mutex should not be poisoned") - .is_empty() - ); - assert!( - notification_methods - .lock() - .expect("notification method mutex should not be poisoned") - .is_empty() - ); - assert_eq!(disconnect_count.load(Ordering::SeqCst), 0); - + assert_eq!(*observed.lock().unwrap(), vec![native_server()]); Ok(()) } #[tokio::test] -async fn unavailable_v2_downstream_rejects_native_declarations() +async fn closing_subscription_stream_cancels_only_its_native_request() -> Result<(), agent_client_protocol::Error> { - let observed = Arc::new(ObservedRequests::default()); - let agent = RecordingAgent { - capabilities: agent_capabilities(v2::McpCapabilities::new()), - observed: Arc::clone(&observed), - }; - let connect_count = Arc::new(AtomicUsize::new(0)); - let request_methods = Arc::new(Mutex::new(Vec::new())); - let notification_methods = Arc::new(Mutex::new(Vec::new())); - let disconnect_count = Arc::new(AtomicUsize::new(0)); - - run_with_polyfill( - agent, - Arc::clone(&connect_count), - Arc::clone(&request_methods), - Arc::clone(¬ification_methods), - Arc::clone(&disconnect_count), - async move |connection| { - let initialize = connection - .send_request(initialize_request()) - .block_task() - .await?; - let mcp = negotiated_mcp_capabilities(&initialize); - assert!(mcp.http.is_none()); - assert!(mcp.acp.is_none()); - - let error = connection - .send_request( - v2::NewSessionRequest::new(PathBuf::from("/tmp")) - .mcp_servers(vec![native_server()]), - ) - .block_task() - .await - .expect_err("native MCP should require a downstream transport"); - assert_eq!(error.code, agent_client_protocol::ErrorCode::InvalidParams); - assert_eq!( - error.data, - Some(serde_json::json!( - "the downstream agent supports neither native nor HTTP MCP transport" - )) + let observed = Arc::new(Mutex::new(Vec::new())); + let cancelled = Arc::new(AtomicUsize::new(0)); + let ids = Arc::new(Mutex::new(Vec::new())); + let caps = v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new() + .mcp(v2::McpCapabilities::new().http(v2::McpHttpCapabilities::new())), + ); + run( + caps, observed.clone(), ids.clone(), Arc::default(), cancelled.clone(), + async |connection| { + connection.send_request(initialize()).block_task().await?; + connection.send_request(v2::NewSessionRequest::new(PathBuf::from("/tmp")) + .mcp_servers(vec![native_server()])).block_task().await?; + let (url, bearer) = { + let observed = observed.lock().unwrap(); + let v2::McpServer::Http(server) = &observed[0] else { panic!("expected HTTP endpoint") }; + (server.url.clone(), server.headers[0].value.clone()) + }; + let address = url.strip_prefix("http://").unwrap(); + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + let body = serde_json::json!({ + "jsonrpc":"2.0","id":73,"method":"subscriptions/listen", + "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28"}} + }).to_string(); + let request = format!( + "POST / HTTP/1.1\r\nHost: {address}\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: subscriptions/listen\r\nContent-Length: {}\r\n\r\n{body}", + body.len() ); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut output = String::new(); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while !output.contains("notifications/subscriptions/acknowledged") { + let mut buf = [0; 2048]; + let n = stream.read(&mut buf).await.unwrap(); + assert_ne!(n, 0, "subscription stream closed before ack"); + output.push_str(std::str::from_utf8(&buf[..n]).unwrap()); + } + }).await.expect("expected subscription acknowledgment"); + assert!(output.contains("\"io.modelcontextprotocol/subscriptionId\":73"), "{output}"); + // A second live POST uses the same external ID, but must retain its + // own generated logical ID and cancellation lifetime. + let mut second = tokio::net::TcpStream::connect(address).await.unwrap(); + second.write_all(request.as_bytes()).await.unwrap(); + let mut second_output = String::new(); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while !second_output.contains("notifications/subscriptions/acknowledged") { + let mut buf = [0; 2048]; + let n = second.read(&mut buf).await.unwrap(); + assert_ne!(n, 0, "second subscription closed before ack"); + second_output.push_str(std::str::from_utf8(&buf[..n]).unwrap()); + } + }).await.expect("expected second subscription acknowledgment"); + assert!(second_output.contains("\"io.modelcontextprotocol/subscriptionId\":73"), "{second_output}"); + let logical_ids = ids.lock().unwrap().clone(); + assert_eq!(logical_ids.len(), 2); + assert_ne!(logical_ids[0], logical_ids[1]); + drop(stream); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while cancelled.load(Ordering::SeqCst) != 1 { + tokio::task::yield_now().await; + } + }).await.expect("closing the HTTP stream must cancel native ACP request"); + drop(second); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while cancelled.load(Ordering::SeqCst) != 2 { + tokio::task::yield_now().await; + } + }).await.expect("closing the second stream must cancel its own ACP request"); + let overflow = post(&url, &bearer, "subscriptions/flood", "").await; + assert_eq!(overflow["error"]["code"], -32000, "{overflow}"); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while cancelled.load(Ordering::SeqCst) != 3 { + tokio::task::yield_now().await; + } + }).await.expect("overflow must cancel only its native ACP request"); + assert_eq!(post(&url, &bearer, "tools/list", "").await["result"]["tools"][0]["name"], "ping"); Ok(()) }, - ) - .await?; - - assert!( - observed - .setup - .lock() - .expect("setup request mutex should not be poisoned") - .is_empty(), - "the rejected request must not reach the downstream agent" - ); - assert_eq!(connect_count.load(Ordering::SeqCst), 0); - assert!( - request_methods - .lock() - .expect("request method mutex should not be poisoned") - .is_empty() - ); - assert!( - notification_methods - .lock() - .expect("notification method mutex should not be poisoned") - .is_empty() - ); - assert_eq!(disconnect_count.load(Ordering::SeqCst), 0); - - Ok(()) + ).await } diff --git a/src/agent-client-protocol-conductor/tests/mcp_server_handler_chain_v2.rs b/src/agent-client-protocol-conductor/tests/mcp_server_handler_chain_v2.rs index 75da7d16..753130df 100644 --- a/src/agent-client-protocol-conductor/tests/mcp_server_handler_chain_v2.rs +++ b/src/agent-client-protocol-conductor/tests/mcp_server_handler_chain_v2.rs @@ -10,8 +10,8 @@ use std::{ }; use agent_client_protocol::{ - Agent, ByteStreams, Client, Conductor, ConnectTo, DynConnectTo, Error, NullRun, Proxy, - Responder, V2ConnectionTo, + Agent, ByteStreams, Client, Conductor, ConnectTo, DynConnectTo, Error, JsonRpcRequest, + JsonRpcResponse, NullRun, Proxy, Responder, V2ConnectionTo, mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, role, schema::{ProtocolVersion, v2}, @@ -22,6 +22,16 @@ use serde_json::json; use tokio::io::duplex; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, JsonRpcRequest)] +#[request(method = "_test/probe", response = ProbeResponse)] +struct ProbeRequest {} + +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, JsonRpcResponse)] +struct ProbeResponse { + #[serde(rename = "resultType")] + result_type: String, +} + fn implementation(name: &str) -> v2::Implementation { v2::Implementation::new(name, env!("CARGO_PKG_VERSION")) } @@ -40,7 +50,7 @@ fn existing_server() -> v2::McpServer { #[derive(Debug, PartialEq, Eq)] struct ObservedMcpContext { server_id: String, - connection_id: String, + request_id: String, } struct RecordingMcpConnect { @@ -58,9 +68,9 @@ impl McpServerConnect for RecordingMcpConnect { .server_id() .expect("the global MCP server should be attached through ACP") .to_string(), - connection_id: context - .connection_id() - .expect("an attached MCP connection should have an ID") + request_id: context + .request_id() + .expect("an attached MCP request should have an ID") .to_string(), }); DynConnectTo::new(PendingMcpComponent) @@ -73,6 +83,14 @@ impl ConnectTo for PendingMcpComponent { async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { role::mcp::Server .builder() + .on_receive_request( + async |_request: ProbeRequest, responder: Responder, _connection| { + responder.respond(ProbeResponse { + result_type: "complete".to_owned(), + }) + }, + agent_client_protocol::on_receive_request!(), + ) .connect_with(client, async |_connection| { std::future::pending::>().await }) @@ -214,14 +232,17 @@ impl ConnectTo for RecordingAgent { let mcp_connection = connection.clone(); connection.spawn(async move { let result = async { - let connected = mcp_connection - .send_request(v2::ConnectMcpRequest::new(server_id)) - .block_task() - .await?; mcp_connection - .send_request(v2::DisconnectMcpRequest::new( - connected.connection_id, - )) + .send_request(v2::MessageMcpRequest::new( + server_id, + v2::McpRequestId::new("global-v2-probe"), + "_test/probe", + ).params(serde_json::Map::from_iter([ + ("_meta".to_owned(), json!({ + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + })), + ]))) .block_task() .await?; Ok(()) @@ -335,7 +356,7 @@ async fn v2_global_mcp_attachment_preserves_setup_and_continues_handler_chain() tokio::time::timeout(std::time::Duration::from_secs(2), round_trip_rx.next()) .await - .expect("global MCP connect/disconnect round trip should not hang") + .expect("global MCP request should not hang") .ok_or_else(|| Error::internal_error().data("MCP round-trip channel closed"))??; connection @@ -369,8 +390,8 @@ async fn v2_global_mcp_attachment_preserves_setup_and_continues_handler_chain() assert_eq!(mcp_contexts.len(), 1); assert_eq!(mcp_contexts[0].server_id, server_ids[0].to_string()); assert!( - !mcp_contexts[0].connection_id.is_empty(), - "the global MCP connection should receive a connection ID" + mcp_contexts[0].request_id == "global-v2-probe", + "the global MCP request should retain its logical request ID" ); Ok(()) } diff --git a/src/agent-client-protocol-conductor/tests/request_cancellation.rs b/src/agent-client-protocol-conductor/tests/request_cancellation.rs index 34eac429..19ca45fc 100644 --- a/src/agent-client-protocol-conductor/tests/request_cancellation.rs +++ b/src/agent-client-protocol-conductor/tests/request_cancellation.rs @@ -21,12 +21,12 @@ use std::time::Duration; use agent_client_protocol::DynConnectTo; use agent_client_protocol::schema::ProtocolVersion; use agent_client_protocol::schema::v1::{ - CancelRequestNotification, ConnectMcpRequest, ContentBlock, ContentChunk, InitializeRequest, - InitializeResponse, McpServer as SchemaMcpServer, McpServerAcpId, NewSessionRequest, - NewSessionResponse, PermissionOption, PermissionOptionKind, PromptRequest, PromptResponse, - RequestPermissionOutcome, RequestPermissionRequest, RequestPermissionResponse, - SelectedPermissionOutcome, SessionId, SessionNotification, SessionUpdate, StopReason, - ToolCallUpdate, ToolCallUpdateFields, + 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, @@ -909,90 +909,101 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - let (probe_barrier_tx, mut probe_barrier_rx) = mpsc::unbounded(); let cancelled_mcp_server_id = Arc::new(Mutex::new(None::)); - let agent = Agent - .builder() - .on_receive_request( - async |initialize: InitializeRequest, responder, _cx: ConnectionTo| { - 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, - cx: ConnectionTo| { + let agent = + Agent + .builder() + .on_receive_request( + async |initialize: InitializeRequest, responder, _cx: ConnectionTo| { + 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(); - let advertised_mcp_server_id = advertised_mcp_server_id(&request); - - if request.cwd.ends_with("park-session") { - *cancelled_mcp_server_id + async move |request: NewSessionRequest, + responder: Responder, + cx: ConnectionTo| { + 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, + >()) + .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") = - Some(advertised_mcp_server_id); - parked_id_tx.unbounded_send(responder.id().clone()).unwrap(); - let cancellation = responder.cancellation(); + .expect("cancelled MCP ID mutex poisoned") + .clone() + .expect("cancelled session should have advertised an MCP server"); + let connection = cx.clone(); cx.spawn(async move { - let response = cancellation - .run_until_cancelled(std::future::pending::< - Result, - >()) - .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(ConnectMcpRequest::new(stale_server_id)) + connection + .send_request(MessageMcpRequest::new( + stale_server_id, + McpRequestId::new("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 |_| 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| { - agent_cancel_tx.unbounded_send(cancel.request_id).unwrap(); - Ok(()) - }, - agent_client_protocol::on_receive_notification!(), - ); + 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| { + 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, @@ -1100,6 +1111,237 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - Ok(()) } +#[derive(Clone)] +struct ParkedMcpServer { + started_tx: mpsc::UnboundedSender, + stopped_tx: mpsc::UnboundedSender, + dropped_tx: mpsc::UnboundedSender<()>, + late_tx: mpsc::UnboundedSender>, +} + +impl McpServerConnect for ParkedMcpServer { + fn name(&self) -> String { + "parked-mcp".into() + } + + fn connect(&self, cx: McpConnectionTo) -> DynConnectTo { + 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 { + sender: mpsc::UnboundedSender, + value: Option, +} + +impl Drop for ProbeOnDrop { + fn drop(&mut self) { + if let Some(value) = self.value.take() { + drop(self.sender.unbounded_send(value)); + } + } +} + +impl ConnectTo for ParkedMcpComponent { + async fn connect_to(self, client: impl ConnectTo) -> 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, + cx: ConnectionTo| { + 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 { + // Request-scoped cancellation drops this whole backend, + // rather than sending a second, inner cancellation RPC. + let _stopped = stopped; + let result = cancellation + .run_until_cancelled(std::future::pending::< + Result, + >()) + .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 {} + +/// An MCP operation retains its logical ID while each ACP transport hop +/// rewrites the outer JSON-RPC ID. Cancelling the ACP request tears down the +/// per-operation server and must not deliver a late MCP notification. +#[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| { + responder.respond(InitializeResponse::new(request.protocol_version)) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: NewSessionRequest, + responder: Responder, + cx: ConnectionTo| { + 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 = 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| { + 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(()) +} + /// `initialize` is rewritten to `_proxy/initialize` at the conductor-to-proxy /// hop and forwarded with a result hook — cancellation must still propagate /// hop by hop, exactly like every other request. diff --git a/src/agent-client-protocol-conductor/tests/scoped_mcp_server.rs b/src/agent-client-protocol-conductor/tests/scoped_mcp_server.rs index d0e3a645..5e3eab5c 100644 --- a/src/agent-client-protocol-conductor/tests/scoped_mcp_server.rs +++ b/src/agent-client-protocol-conductor/tests/scoped_mcp_server.rs @@ -39,7 +39,7 @@ async fn test_scoped_mcp_server_through_proxy() -> Result<(), agent_client_proto .await?; expect_test::expect![[r#" - "OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"2\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" + "OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"2\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" "#]].assert_debug_eq(&result); Ok(()) @@ -84,7 +84,7 @@ async fn test_scoped_mcp_server_through_session() -> Result<(), agent_client_pro .await?; expect_test::expect![[r#" - "OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"2\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" + "OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"2\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" "#]].assert_debug_eq(&result); Ok(()) diff --git a/src/agent-client-protocol-conductor/tests/standalone_mcp_server.rs b/src/agent-client-protocol-conductor/tests/standalone_mcp_server.rs index 8de40f55..33118cb1 100644 --- a/src/agent-client-protocol-conductor/tests/standalone_mcp_server.rs +++ b/src/agent-client-protocol-conductor/tests/standalone_mcp_server.rs @@ -41,7 +41,7 @@ fn create_test_server() -> McpServer DynConnectTo { @@ -28,15 +27,15 @@ fn create_echo_proxy() -> DynConnectTo { .instructions("Test MCP server with a connection-context echo tool") .tool_fn_mut( "echo", - "Returns the current MCP connection context", + "Returns the current MCP request context", async |_input: EchoInput, context| { Ok(EchoOutput { server_id: context .server_id() .expect("tool is attached through ACP") .to_string(), - connection_id: context - .connection_id() + request_id: context + .request_id() .expect("tool is attached through ACP") .to_string(), }) @@ -88,7 +87,7 @@ async fn test_list_tools_from_mcp_server() -> Result<(), agent_client_protocol:: expect![[r" Available tools: - - echo: Returns the current MCP connection context"]] + - echo: Returns the current MCP request context"]] .assert_eq(&result); Ok(()) @@ -115,14 +114,10 @@ async fn test_acp_identifiers_are_delivered_to_mcp_tools() let server_id = regex::Regex::new(r#""server_id":\s*String\("mcp-server:[0-9a-f-]+"\)"#) .expect("valid server ID regex"); - let connection_id = - regex::Regex::new(r#""connection_id":\s*String\("mcp-over-acp-connection:[0-9a-f-]+"\)"#) - .expect("valid connection ID regex"); + let request_id = regex::Regex::new(r#""request_id":\s*String\("[^"]+"\)"#) + .expect("valid logical request ID regex"); assert!(server_id.is_match(&result), "unexpected result: {result}"); - assert!( - connection_id.is_match(&result), - "unexpected result: {result}" - ); + assert!(request_id.is_match(&result), "unexpected result: {result}"); Ok(()) } diff --git a/src/agent-client-protocol-conductor/tests/test_tool_fn.rs b/src/agent-client-protocol-conductor/tests/test_tool_fn.rs index 55203413..67f91609 100644 --- a/src/agent-client-protocol-conductor/tests/test_tool_fn.rs +++ b/src/agent-client-protocol-conductor/tests/test_tool_fn.rs @@ -74,7 +74,7 @@ async fn test_tool_fn_greet() -> Result<(), agent_client_protocol::Error> { .await?; expect_test::expect![[r#" - "OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"\\\"Hello, World!\\\"\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" + "OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"\\\"Hello, World!\\\"\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" "#]].assert_debug_eq(&result); Ok(()) diff --git a/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs b/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs index 3987983a..11c51fed 100644 --- a/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs +++ b/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs @@ -33,7 +33,8 @@ use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; /// - Replaces UUIDs with sequential IDs (id:0, id:1, etc.) /// - Replaces session IDs with "session:0", etc. /// - Replaces loopback HTTP endpoints with "http:endpoint:0", etc. -/// - Replaces MCP server and connection IDs with stable sequential IDs +/// - Replaces MCP server and logical request IDs with stable sequential IDs +/// - Redacts runtime-generated HTTP Authorization headers struct EventNormalizer { id_map: HashMap, next_id: usize, @@ -43,8 +44,8 @@ struct EventNormalizer { next_endpoint: usize, server_map: HashMap, next_server: usize, - connection_map: HashMap, - next_connection: usize, + request_map: HashMap, + next_request: usize, } impl EventNormalizer { @@ -58,8 +59,8 @@ impl EventNormalizer { next_endpoint: 0, server_map: HashMap::new(), next_server: 0, - connection_map: HashMap::new(), - next_connection: 0, + request_map: HashMap::new(), + next_request: 0, } } @@ -116,12 +117,12 @@ impl EventNormalizer { .clone() } - fn normalize_connection_id(&mut self, id: &str) -> String { - self.connection_map + fn normalize_request_id(&mut self, id: &str) -> String { + self.request_map .entry(id.to_string()) .or_insert_with(|| { - let n = format!("connection:{}", self.next_connection); - self.next_connection += 1; + let n = format!("request:{}", self.next_request); + self.next_request += 1; n }) .clone() @@ -131,6 +132,10 @@ impl EventNormalizer { fn normalize_json(&mut self, value: serde_json::Value) -> serde_json::Value { match value { serde_json::Value::Object(map) => { + let is_authorization_header = map + .get("name") + .and_then(|v| v.as_str()) + .is_some_and(|name| name.eq_ignore_ascii_case("authorization")); let normalized: serde_json::Map = map .into_iter() .map(|(k, v)| { @@ -158,12 +163,14 @@ impl EventNormalizer { } else { self.normalize_json(v) } - } else if k == "connectionId" { + } else if k == "requestId" { if let serde_json::Value::String(s) = &v { - serde_json::Value::String(self.normalize_connection_id(s)) + serde_json::Value::String(self.normalize_request_id(s)) } else { self.normalize_json(v) } + } else if is_authorization_header && k == "value" { + serde_json::Value::String("[REDACTED]".into()) } else { self.normalize_json(v) }; @@ -500,7 +507,7 @@ async fn test_trace_client_mcp_server() -> Result<(), agent_client_protocol::Err "sessionUpdate": String("agent_message_chunk"), "content": Object { "type": String("text"), - "text": String("OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"{\\\"echoed\\\":\\\"Client echoes: Hello from client test!\\\",\\\"call_number\\\":1}\", meta: None, annotations: None })], structured_content: Some(Object {\"echoed\": String(\"Client echoes: Hello from client test!\"), \"call_number\": Number(1)}), is_error: Some(false), meta: None }"), + "text": String("OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"{\\\"echoed\\\":\\\"Client echoes: Hello from client test!\\\",\\\"call_number\\\":1}\", meta: None, annotations: None })], structured_content: Some(Object {\"echoed\": String(\"Client echoes: Hello from client test!\"), \"call_number\": Number(1)}), is_error: Some(false), meta: None }"), }, "messageId": String("testy-message-end-turn-1"), }, diff --git a/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs b/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs index 48366ebf..f8318573 100644 --- a/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs +++ b/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs @@ -32,7 +32,8 @@ use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; /// - Replaces UUIDs with sequential IDs (id:0, id:1, etc.) /// - Replaces session IDs with "session:0", etc. /// - Replaces loopback HTTP endpoints with "http:endpoint:0", etc. -/// - Replaces MCP server and connection IDs with stable sequential IDs +/// - Replaces MCP server and logical request IDs with stable sequential IDs +/// - Redacts runtime-generated HTTP Authorization headers struct EventNormalizer { id_map: HashMap, next_id: usize, @@ -42,8 +43,8 @@ struct EventNormalizer { next_endpoint: usize, server_map: HashMap, next_server: usize, - connection_map: HashMap, - next_connection: usize, + request_map: HashMap, + next_request: usize, } impl EventNormalizer { @@ -57,8 +58,8 @@ impl EventNormalizer { next_endpoint: 0, server_map: HashMap::new(), next_server: 0, - connection_map: HashMap::new(), - next_connection: 0, + request_map: HashMap::new(), + next_request: 0, } } @@ -115,12 +116,12 @@ impl EventNormalizer { .clone() } - fn normalize_connection_id(&mut self, id: &str) -> String { - self.connection_map + fn normalize_request_id(&mut self, id: &str) -> String { + self.request_map .entry(id.to_string()) .or_insert_with(|| { - let n = format!("connection:{}", self.next_connection); - self.next_connection += 1; + let n = format!("request:{}", self.next_request); + self.next_request += 1; n }) .clone() @@ -130,6 +131,10 @@ impl EventNormalizer { fn normalize_json(&mut self, value: serde_json::Value) -> serde_json::Value { match value { serde_json::Value::Object(map) => { + let is_authorization_header = map + .get("name") + .and_then(|v| v.as_str()) + .is_some_and(|name| name.eq_ignore_ascii_case("authorization")); let normalized: serde_json::Map = map .into_iter() .map(|(k, v)| { @@ -157,12 +162,14 @@ impl EventNormalizer { } else { self.normalize_json(v) } - } else if k == "connectionId" { + } else if k == "requestId" { if let serde_json::Value::String(s) = &v { - serde_json::Value::String(self.normalize_connection_id(s)) + serde_json::Value::String(self.normalize_request_id(s)) } else { self.normalize_json(v) } + } else if is_authorization_header && k == "value" { + serde_json::Value::String("[REDACTED]".into()) } else { self.normalize_json(v) }; @@ -317,7 +324,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> // Snapshot the trace events // This should show: // 1. Client -> Agent: initialize, session/new, session/prompt (left-to-right) - // 2. Agent -> MCP Server: tools/call (right-to-left, the key part!) + // 2. Agent -> MCP Server: discovery/list/call with per-request MCP metadata // 3. MCP Server -> Agent: response // 4. Agent -> Client: notification + response expect![[r#" @@ -490,32 +497,6 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> }, }, ), - Request( - RequestEvent { - ts: 0.0, - protocol: Acp, - from: "Proxy(1)", - to: "Proxy(0)", - id: String("id:4"), - method: "mcp/connect", - session: None, - params: Object { - "serverId": String("server:0"), - }, - }, - ), - Response( - ResponseEvent { - ts: 0.0, - from: "Proxy(0)", - to: "Proxy(1)", - id: String("id:4"), - is_error: false, - payload: Object { - "connectionId": String("connection:0"), - }, - }, - ), Response( ResponseEvent { ts: 0.0, @@ -622,7 +603,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> protocol: Acp, from: "Client", to: "Proxy(0)", - id: String("id:5"), + id: String("id:4"), method: "session/prompt", session: None, params: Object { @@ -642,7 +623,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> protocol: Acp, from: "Proxy(0)", to: "Proxy(1)", - id: String("id:6"), + id: String("id:5"), method: "session/prompt", session: None, params: Object { @@ -662,15 +643,17 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> protocol: Mcp, from: "Proxy(1)", to: "Proxy(0)", - id: String("id:7"), - method: "initialize", + id: String("id:6"), + method: "server/discover", session: None, params: Object { - "protocolVersion": String("2025-11-25"), - "capabilities": Object {}, - "clientInfo": Object { - "name": String("rmcp"), - "version": String("3.4.0"), + "_meta": Object { + "io.modelcontextprotocol/protocolVersion": String("2026-07-28"), + "io.modelcontextprotocol/clientInfo": Object { + "name": String("testy"), + "version": String("0.11.0"), + }, + "io.modelcontextprotocol/clientCapabilities": Object {}, }, }, }, @@ -680,30 +663,98 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> ts: 0.0, from: "Proxy(0)", to: "Proxy(1)", - id: String("id:7"), + id: String("id:6"), is_error: false, payload: Object { - "protocolVersion": String("2025-11-25"), + "resultType": String("complete"), + "supportedVersions": Array [ + String("2026-07-28"), + ], "capabilities": Object { "tools": Object {}, }, - "serverInfo": Object { - "name": String("rmcp"), - "version": String("3.4.0"), - }, "instructions": String("A simple test MCP server with an echo tool"), + "ttlMs": Number(0), + "cacheScope": String("private"), + "_meta": Object { + "io.modelcontextprotocol/serverInfo": Object { + "name": String("rmcp"), + "version": String("3.4.0"), + }, + }, }, }, ), - Notification( - NotificationEvent { + Request( + RequestEvent { ts: 0.0, protocol: Mcp, from: "Proxy(1)", to: "Proxy(0)", - method: "notifications/initialized", + id: String("id:7"), + method: "tools/list", session: None, - params: Null, + params: Object { + "_meta": Object { + "io.modelcontextprotocol/protocolVersion": String("2026-07-28"), + "io.modelcontextprotocol/clientInfo": Object { + "name": String("testy"), + "version": String("0.11.0"), + }, + "io.modelcontextprotocol/clientCapabilities": Object {}, + "progressToken": Number(0), + }, + }, + }, + ), + Response( + ResponseEvent { + ts: 0.0, + from: "Proxy(0)", + to: "Proxy(1)", + id: String("id:7"), + is_error: false, + payload: Object { + "resultType": String("complete"), + "ttlMs": Number(0), + "cacheScope": String("private"), + "tools": Array [ + Object { + "name": String("echo"), + "description": String("Echoes back the input message"), + "inputSchema": Object { + "$schema": String("https://json-schema.org/draft/2020-12/schema"), + "title": String("EchoParams"), + "description": String("Parameters for the echo tool"), + "type": String("object"), + "properties": Object { + "message": Object { + "description": String("The message to echo back"), + "type": String("string"), + }, + }, + "required": Array [ + String("message"), + ], + }, + "outputSchema": Object { + "$schema": String("https://json-schema.org/draft/2020-12/schema"), + "title": String("EchoOutput"), + "description": String("Output from the echo tool"), + "type": String("object"), + "properties": Object { + "result": Object { + "description": String("The echoed message"), + "type": String("string"), + }, + }, + "required": Array [ + String("result"), + ], + }, + }, + ], + }, }, ), Request( @@ -717,6 +768,12 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> session: None, params: Object { "_meta": Object { + "io.modelcontextprotocol/protocolVersion": String("2026-07-28"), + "io.modelcontextprotocol/clientInfo": Object { + "name": String("testy"), + "version": String("0.11.0"), + }, + "io.modelcontextprotocol/clientCapabilities": Object {}, "progressToken": Number(0), }, "name": String("echo"), @@ -734,6 +791,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> id: String("id:8"), is_error: false, payload: Object { + "resultType": String("complete"), "content": Array [ Object { "type": String("text"), @@ -761,7 +819,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> "sessionUpdate": String("agent_message_chunk"), "content": Object { "type": String("text"), - "text": String("OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"{\\\"result\\\":\\\"Echo: Hello from trace test!\\\"}\", meta: None, annotations: None })], structured_content: Some(Object {\"result\": String(\"Echo: Hello from trace test!\")}), is_error: Some(false), meta: None }"), + "text": String("OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"{\\\"result\\\":\\\"Echo: Hello from trace test!\\\"}\", meta: None, annotations: None })], structured_content: Some(Object {\"result\": String(\"Echo: Hello from trace test!\")}), is_error: Some(false), meta: None }"), }, "messageId": String("testy-message-end-turn-1"), }, @@ -773,7 +831,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> ts: 0.0, from: "Proxy(1)", to: "Proxy(0)", - id: String("id:6"), + id: String("id:5"), is_error: false, payload: Object { "stopReason": String("end_turn"), @@ -794,7 +852,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> "sessionUpdate": String("agent_message_chunk"), "content": Object { "type": String("text"), - "text": String("OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"{\\\"result\\\":\\\"Echo: Hello from trace test!\\\"}\", meta: None, annotations: None })], structured_content: Some(Object {\"result\": String(\"Echo: Hello from trace test!\")}), is_error: Some(false), meta: None }"), + "text": String("OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"{\\\"result\\\":\\\"Echo: Hello from trace test!\\\"}\", meta: None, annotations: None })], structured_content: Some(Object {\"result\": String(\"Echo: Hello from trace test!\")}), is_error: Some(false), meta: None }"), }, "messageId": String("testy-message-end-turn-1"), }, @@ -806,7 +864,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> ts: 0.0, from: "Proxy(0)", to: "Client", - id: String("id:5"), + id: String("id:4"), is_error: false, payload: Object { "stopReason": String("end_turn"), diff --git a/src/agent-client-protocol-cookbook/src/lib.rs b/src/agent-client-protocol-cookbook/src/lib.rs index 58fba618..e5686a45 100644 --- a/src/agent-client-protocol-cookbook/src/lib.rs +++ b/src/agent-client-protocol-cookbook/src/lib.rs @@ -732,7 +732,7 @@ pub mod global_mcp_server { //! ``` //! //! The `from_rmcp` function takes a factory closure that creates a new server - //! instance. This allows each MCP connection to get a fresh server instance. + //! instance for each MCP request. //! //! # How it works //! @@ -740,13 +740,14 @@ pub mod global_mcp_server { //! handler. It: //! //! 1. Intercepts session setup requests and adds a schema-native - //! `McpServer::Acp` declaration with one connection-scoped server ID. + //! `McpServer::Acp` declaration with one stable server ID. //! V1 injects it into `session/new`, `session/load`, `session/resume`, //! and feature-gated `session/fork`; v2 injects it into //! `session/new`, `session/resume`, and feature-gated `session/fork` //! while preserving unrelated request fields //! 2. Passes the modified request through to the next handler - //! 3. Handles `mcp/connect`, `mcp/message`, and `mcp/disconnect` for that server ID + //! 3. Handles `mcp/message` requests for that server ID. Each operation + //! has its own logical request ID and per-request MCP metadata. //! //! [`McpServer::builder`]: agent_client_protocol_rmcp::McpServerExt::builder //! [`McpServer::from_rmcp`]: agent_client_protocol_rmcp::McpServerExt::from_rmcp diff --git a/src/agent-client-protocol-polyfill/CHANGELOG.md b/src/agent-client-protocol-polyfill/CHANGELOG.md index 77cb9487..c02f4ee7 100644 --- a/src/agent-client-protocol-polyfill/CHANGELOG.md +++ b/src/agent-client-protocol-polyfill/CHANGELOG.md @@ -7,6 +7,21 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Changed + +- Replace the MCP session bridge with a latest-only MCP 2026-07-28 HTTP adapter: + one native ACP request per POST, request-scoped SSE, and stream-close + cancellation. Remove the initialize/connect/disconnect, GET, batch, and + session-header paths; existing older MCP HTTP clients must be upgraded. +- Require runtime bearer credentials from the rewritten declaration, validate + Origin and mirrored routing headers, and preserve logical request identities + independently of overlapping HTTP IDs. +- Bound notification queues and request/listener admission. Translate + subscription IDs in namespaced MCP metadata for notifications and completion. +- Fail closed for unsupported `x-mcp-header` tools. Direct tool calls perform + descriptor lookup internally rather than requiring a client-side list + handshake. This remains a draft, not a claim of full HTTP conformance. + ## [2.2.0](https://github.com/agentclientprotocol/rust-sdk/compare/agent-client-protocol-polyfill-v2.1.0...agent-client-protocol-polyfill-v2.2.0) - 2026-09-18 ### Other diff --git a/src/agent-client-protocol-polyfill/Cargo.toml b/src/agent-client-protocol-polyfill/Cargo.toml index 6c678b4f..0cb0d67c 100644 --- a/src/agent-client-protocol-polyfill/Cargo.toml +++ b/src/agent-client-protocol-polyfill/Cargo.toml @@ -19,14 +19,15 @@ unstable_session_fork = ["agent-client-protocol/unstable_session_fork"] agent-client-protocol = { workspace = true, features = ["unstable_mcp_over_acp"] } async-stream.workspace = true axum.workspace = true +base64.workspace = true futures.workspace = true -futures-concurrency.workspace = true -rustc-hash.workspace = true serde_json.workspace = true -thiserror = "2.0" tokio = { workspace = true, features = ["net"] } tracing.workspace = true uuid.workspace = true [lints] workspace = true + +[dev-dependencies] +tokio = { workspace = true, features = ["io-util", "macros", "rt"] } diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs deleted file mode 100644 index 908b6b87..00000000 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs +++ /dev/null @@ -1,78 +0,0 @@ -use agent_client_protocol::{ConnectTo, Dispatch, DynConnectTo, role::mcp}; -use futures::{SinkExt as _, StreamExt as _, channel::mpsc}; -use tracing::info; - -use super::BridgeMessage; - -/// Actor that bridges a single MCP connection between a local MCP client -/// and the ACP proxy chain. -#[derive(Debug)] -pub(crate) struct BridgeConnectionActor { - /// The loopback HTTP transport accepted by the compatibility listener. - transport: DynConnectTo, - - /// Sender for messages back to the polyfill's bridge runner loop. - bridge_tx: mpsc::Sender, - - /// Receiver for messages from the polyfill to forward to the MCP client. - to_mcp_client_rx: mpsc::Receiver, -} - -impl BridgeConnectionActor { - pub fn new( - component: impl ConnectTo, - bridge_tx: mpsc::Sender, - to_mcp_client_rx: mpsc::Receiver, - ) -> Self { - Self { - transport: DynConnectTo::new(component), - bridge_tx, - to_mcp_client_rx, - } - } - - pub async fn run(self, connection_id: String) -> Result<(), agent_client_protocol::Error> { - info!(connection_id, "MCP bridge connected"); - - let Self { - transport, - mut bridge_tx, - to_mcp_client_rx, - } = self; - - let result = mcp::Client - .builder() - .name(format!("mcp-client-to-polyfill({connection_id})")) - .on_receive_dispatch( - { - let mut bridge_tx = bridge_tx.clone(); - let connection_id = connection_id.clone(); - async move |message: Dispatch, _cx| { - bridge_tx - .send(BridgeMessage::ClientToServer { - connection_id: connection_id.clone(), - message, - }) - .await - .map_err(|_| agent_client_protocol::Error::internal_error()) - } - }, - agent_client_protocol::on_receive_dispatch!(), - ) - .connect_with(transport, async move |mcp_connection_to_client| { - let mut to_mcp_client_rx = to_mcp_client_rx; - while let Some(message) = to_mcp_client_rx.next().await { - mcp_connection_to_client.send_proxied_message(message)?; - } - Ok(()) - }) - .await; - - bridge_tx - .send(BridgeMessage::Disconnected { connection_id }) - .await - .map_err(|_| agent_client_protocol::Error::internal_error())?; - - result - } -} diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs index afed504b..9389a81a 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs @@ -1,1402 +1,540 @@ -//! HTTP-based MCP bridge transport. +//! MCP 2026-07-28 request-scoped Streamable HTTP endpoint. +//! +//! The adapter keeps raw MCP envelopes, ACP cancellation, and response-stream +//! lifetimes explicit. Each POST owns one operation, not an MCP session. -use agent_client_protocol::{ - BoxFuture, Channel, ConnectTo, RawJsonRpcMessage, RawJsonRpcParams, TransportBatchEntry, - TransportFrame, - role::mcp, - schema::v1::{ - Notification as RpcNotification, Request as RpcRequest, RequestId, Response as RpcResponse, - }, -}; +use std::{convert::Infallible, sync::Arc}; + +use agent_client_protocol::Error; use axum::{ - Router, + Json, Router, + body::Bytes, extract::State, - http::StatusCode, - response::{IntoResponse, Response, Sse}, + http::{HeaderMap, StatusCode, header}, + response::{ + IntoResponse, Response, Sse, + sse::{Event, KeepAlive}, + }, routing::post, }; -use futures::{SinkExt, StreamExt as _, channel::mpsc, future::Either, stream::Stream}; -use futures_concurrency::future::FutureExt as _; -use futures_concurrency::stream::StreamExt as _; -use rustc_hash::FxHashMap; -use std::{ - collections::{HashMap, VecDeque}, - pin::pin, - sync::Arc, +use base64::Engine as _; +use futures::{SinkExt, channel::mpsc}; +use serde_json::{Map, Value}; +use tokio::{ + net::TcpListener, + sync::{mpsc as tokio_mpsc, oneshot}, }; -use tokio::net::TcpListener; -use super::{BridgeConnection, BridgeMessage, actor::BridgeConnectionActor}; +use super::BridgeMessage; -/// Runs an HTTP listener for MCP bridge connections. -pub async fn run_http_listener( - tcp_listener: TcpListener, +const VERSION: &str = "2026-07-28"; + +struct BridgeState { server_id: String, - mut bridge_tx: mpsc::Sender, -) -> Result<(), agent_client_protocol::Error> { - let (to_mcp_client_tx, to_mcp_client_rx) = mpsc::channel(128); + token: String, + tx: mpsc::Sender, +} - bridge_tx - .send(BridgeMessage::ConnectionReceived { - server_id, - actor: BridgeConnectionActor::new( - HttpMcpBridge::new(tcp_listener), - bridge_tx.clone(), - to_mcp_client_rx, - ), - connection: BridgeConnection::new(to_mcp_client_tx), - }) +pub(super) async fn run_http_listener( + listener: TcpListener, + server_id: String, + token: String, + tx: mpsc::Sender, +) -> Result<(), Error> { + let state = Arc::new(BridgeState { + server_id, + token, + tx, + }); + let app = Router::new() + .route("/", post(handle_post)) + .with_state(state); + axum::serve(listener, app) .await - .map_err(|_| agent_client_protocol::Error::internal_error())?; - - Ok(()) + .map_err(Error::into_internal_error) } -/// A component that receives HTTP requests/responses using the HTTP transport -/// defined by the MCP protocol. -struct HttpMcpBridge { - listener: tokio::net::TcpListener, +fn error(status: StatusCode, id: Value, code: i64, message: &str) -> Response { + (status, Json(rpc_error(id, code, message))).into_response() } -impl HttpMcpBridge { - /// Creates a new HTTP-MCP bridge from an existing TCP listener. - fn new(listener: tokio::net::TcpListener) -> Self { - Self { listener } +pub(super) fn rpc_error(id: Value, code: i64, message: &str) -> Value { + let mut response = + serde_json::json!({"jsonrpc":"2.0", "error":{"code":code,"message":message}}); + if !id.is_null() { + response["id"] = id; } + response } -impl ConnectTo for HttpMcpBridge { - async fn connect_to( - self, - client: impl ConnectTo, - ) -> Result<(), agent_client_protocol::Error> { - let (channel, serve_self) = self.into_channel_and_future(); - match futures::future::select(pin!(client.connect_to(channel)), serve_self).await { - Either::Left((result, _)) | Either::Right((result, _)) => result, - } - } +fn valid_request_id(id: &Value) -> bool { + id.is_string() || id.as_i64().is_some() || id.as_u64().is_some() +} - fn into_channel_and_future( - self, - ) -> ( - Channel, - BoxFuture<'static, Result<(), agent_client_protocol::Error>>, - ) - where - Self: Sized, +/// Only the MCP 2026 payload metadata carries a subscription identifier. +/// Other fields (including opaque requestState and progress tokens) are untouched. +pub(super) fn rewrite_subscription_id(payload: &mut Value, request_id: &str, http_id: &Value) { + if let Some(subscription_id) = payload + .get_mut("_meta") + .and_then(Value::as_object_mut) + .and_then(|meta| meta.get_mut("io.modelcontextprotocol/subscriptionId")) + && subscription_id.as_str() == Some(request_id) { - let (channel_a, channel_b) = Channel::duplex(); - (channel_a, Box::pin(run(self.listener, channel_b))) + *subscription_id = http_id.clone(); } } -/// Error type for responding to malformed HTTP requests. -#[derive(Debug, thiserror::Error)] -#[error(transparent)] -struct HttpError(#[from] agent_client_protocol::Error); - -impl From for HttpError { - fn from(error: axum::Error) -> Self { - HttpError(agent_client_protocol::util::internal_error(error)) - } +pub(super) fn rpc_result(id: Value, request_id: &str, mut result: Value) -> Value { + rewrite_subscription_id(&mut result, request_id, &id); + serde_json::json!({"jsonrpc":"2.0", "id":id, "result":result}) } -impl IntoResponse for HttpError { - fn into_response(self) -> Response { - let message = format!("Error: {}", self.0); - (StatusCode::INTERNAL_SERVER_ERROR, message).into_response() - } +pub(super) fn rpc_acp_error(id: Value, error: Error) -> Value { + serde_json::json!({"jsonrpc":"2.0", "id":id, "error":error}) } -/// Run a webserver listening on `listener` for HTTP requests at `/` -/// and communicating those requests over `channel` to the JSON-RPC server. -async fn run(listener: TcpListener, channel: Channel) -> Result<(), agent_client_protocol::Error> { - let (registration_tx, registration_rx) = mpsc::unbounded(); - let state = BridgeState { registration_tx }; - - // The way that the MCP protocol works is a bit "special". - // - // Clients *POST* messages to `/`. Those are submitted to the MCP server. - // If the message is a REQUEST, then the client waits until it gets a reply. - // It expects the server to close the connection after responding. - // - // Clients can also issue a *GET* request. This will result in a stream of messages. - // - // Non-reply messages can be sent to any open stream (POST, GET, etc) but must be sent to - // exactly one. - // - // There are provisions for "resuming" from a blocked point by tagging each message in the SSE - // stream with an id, but we are not implementing that because I am lazy. - async { - let app = Router::new() - .route("/", post(handle_post).get(handle_get)) - .with_state(Arc::new(state)); - - axum::serve(listener, app) - .await - .map_err(agent_client_protocol::util::internal_error) - } - .race(RunningServer::new().run(channel, registration_rx)) - .await +fn header_value<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> { + let mut values = headers.get_all(name).iter(); + let value = values.next()?.to_str().ok()?; + values.next().is_none().then_some(value) } -/// The state we pass to our POST/GET handlers. -struct BridgeState { - /// Where to send registration messages. - registration_tx: mpsc::UnboundedSender, +fn valid_origin(headers: &HeaderMap) -> bool { + // Browsers supply Origin; only the actual loopback origin is trusted. + // Non-browser HTTP clients normally omit Origin. + let Some(origin) = header_value(headers, "origin") else { + return !headers.contains_key("origin"); + }; + let Some(host) = header_value(headers, "host") else { + return false; + }; + host.split_once(':') + .is_some_and(|(address, port)| address == "127.0.0.1" && port.parse::().is_ok()) + && origin == format!("http://{host}") } -/// Messages from HTTP handlers to the bridge server. -#[derive(Debug)] -#[allow(dead_code)] -enum HttpMessage { - /// A JSON-RPC request (has an id, expects a response via the channel). - Request { - http_request_id: uuid::Uuid, - request: RpcRequest, - response_tx: mpsc::UnboundedSender, - }, - /// A JSON-RPC notification (no id, no response expected). - Notification { - http_request_id: uuid::Uuid, - request: RpcNotification, - }, - /// A JSON-RPC response from the client. - Response { - http_request_id: uuid::Uuid, - response: RpcResponse, - }, - /// A batch retained as one transport frame. - Frame { - http_request_id: uuid::Uuid, - frame: TransportFrame, - request_ids: Vec, - response_tx: Option>, - }, - /// A GET request to open an SSE stream for server-initiated messages. - Get { - http_request_id: uuid::Uuid, - response_tx: mpsc::UnboundedSender, - }, +fn accepts_both(headers: &HeaderMap) -> bool { + let Some(accept) = header_value(headers, "accept") else { + return false; + }; + let types = accept + .split(',') + .map(|part| part.split(';').next().unwrap_or("").trim()); + let types: Vec<_> = types.collect(); + types.contains(&"application/json") && types.contains(&"text/event-stream") } -struct RunningServer { - waiting_sessions: FxHashMap, - waiting_batch_sessions: Vec, - pending_calls: VecDeque, - general_sessions: Vec, - message_deque: VecDeque, +fn mirrored_name<'a>(method: &str, params: &'a Map) -> Option<&'a str> { + match method { + "tools/call" | "prompts/get" => params.get("name").and_then(Value::as_str), + "resources/read" => params.get("uri").and_then(Value::as_str), + _ => None, + } } -impl RunningServer { - fn new() -> Self { - RunningServer { - waiting_sessions: HashMap::default(), - waiting_batch_sessions: Vec::new(), - pending_calls: VecDeque::new(), - general_sessions: Vec::default(), - message_deque: VecDeque::with_capacity(32), - } +/// Decode the MCP sentinel; rejecting invalid or noncanonical Base64 prevents +/// intermediaries and the adapter from disagreeing on mirrored routing values. +fn matches_mirror(header: Option<&str>, body: &str) -> bool { + let Some(header) = header else { + return false; + }; + if let Some(encoded) = header + .strip_prefix("=?base64?") + .and_then(|h| h.strip_suffix("?=")) + { + base64::engine::general_purpose::STANDARD + .decode(encoded) + .is_ok_and(|bytes| bytes == body.as_bytes()) + } else { + // Literal sentinel-looking values must be encoded to avoid ambiguity. + !header.starts_with("=?base64?") && header == body } +} - /// The main loop: listen for incoming HTTP messages and outgoing JSON-RPC messages. - async fn run( - mut self, - mut channel: Channel, - http_rx: mpsc::UnboundedReceiver, - ) -> Result<(), agent_client_protocol::Error> { - #[derive(Debug)] - enum MultiplexMessage { - FromHttpToChannel(HttpMessage), - FromChannelToHttp(TransportFrame), - } - - let mut merged_stream = http_rx - .map(MultiplexMessage::FromHttpToChannel) - .merge(channel.rx.map(MultiplexMessage::FromChannelToHttp)); - - while let Some(message) = merged_stream.next().await { - tracing::trace!(?message, "received message"); - - match message { - MultiplexMessage::FromHttpToChannel(http_message) => { - self.handle_http_message(http_message, &mut channel.tx)?; - } - MultiplexMessage::FromChannelToHttp(message) => { - self.message_deque.push_back(message); - } - } - - self.drain_jsonrpc_messages(); - self.activate_pending_calls(&mut channel.tx)?; - } - - Ok(()) +async fn handle_post( + State(state): State>, + headers: HeaderMap, + body: Bytes, +) -> Response { + if [ + "origin", + "authorization", + "mcp-protocol-version", + "mcp-method", + "mcp-name", + ] + .into_iter() + .any(|name| headers.get_all(name).iter().nth(1).is_some()) + { + return error( + StatusCode::BAD_REQUEST, + Value::Null, + -32020, + "HeaderMismatch: duplicate routing or authentication header", + ); } - - /// Handle an incoming HTTP message (request, notification, response, or GET). - fn handle_http_message( - &mut self, - message: HttpMessage, - channel_tx: &mut mpsc::UnboundedSender, - ) -> Result<(), agent_client_protocol::Error> { - match message { - HttpMessage::Request { - http_request_id, - request, - response_tx, - } => { - tracing::debug!(%http_request_id, ?request, "handling request"); - let request_id = request.id.clone(); - self.send_or_queue_call( - PendingCall { - frame: TransportFrame::Single(RawJsonRpcMessage::Request(request)), - request_ids: vec![request_id], - session: RegisteredSession::new(response_tx), - }, - channel_tx, - )?; - } - HttpMessage::Notification { - http_request_id: _, - request, - } => { - channel_tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::Notification( - request, - ))) - .map_err(agent_client_protocol::util::internal_error)?; - } - HttpMessage::Response { - http_request_id: _, - response, - } => { - channel_tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::Response( - response, - ))) - .map_err(agent_client_protocol::util::internal_error)?; - } - HttpMessage::Frame { - http_request_id, - frame, - request_ids, - response_tx, - } => { - tracing::debug!(%http_request_id, ?frame, "handling retained frame"); - if let Some(response_tx) = response_tx { - match &frame { - TransportFrame::Batch(_) => { - self.send_or_queue_call( - PendingCall { - frame, - request_ids, - session: RegisteredSession::new(response_tx), - }, - channel_tx, - )?; - return Ok(()); - } - TransportFrame::Single(_) | TransportFrame::Malformed { .. } => { - unreachable!("only batches use the retained frame variant") - } - } - } - channel_tx - .unbounded_send(frame) - .map_err(agent_client_protocol::util::internal_error)?; - } - HttpMessage::Get { - http_request_id: _, - response_tx, - } => { - self.general_sessions - .push(RegisteredSession::new(response_tx)); - } - } - self.purge_closed_sessions(); - Ok(()) + if !valid_origin(&headers) { + return error(StatusCode::FORBIDDEN, Value::Null, -32600, "Invalid Origin"); } - - fn send_or_queue_call( - &mut self, - call: PendingCall, - channel_tx: &mut mpsc::UnboundedSender, - ) -> Result<(), agent_client_protocol::Error> { - if self.call_conflicts_with_active(&call.request_ids) { - tracing::debug!( - request_ids = ?call.request_ids, - "queueing HTTP call until overlapping request IDs are no longer in flight" - ); - self.pending_calls.push_back(call); - return Ok(()); - } - - self.activate_call(call, channel_tx) + if header_value(&headers, "authorization") != Some(&format!("Bearer {}", state.token)) { + return error( + StatusCode::UNAUTHORIZED, + Value::Null, + -32600, + "Unauthorized", + ); } - - fn activate_call( - &mut self, - call: PendingCall, - channel_tx: &mut mpsc::UnboundedSender, - ) -> Result<(), agent_client_protocol::Error> { - let PendingCall { - frame, - request_ids, - session, - } = call; - let is_batch = matches!(frame, TransportFrame::Batch(_)); - channel_tx - .unbounded_send(frame) - .map_err(agent_client_protocol::util::internal_error)?; - - if is_batch { - self.waiting_batch_sessions.push(WaitingBatchSession { - request_ids, - session, - }); - } else { - let request_id = request_ids - .into_iter() - .next() - .expect("single request calls always have one request ID"); - self.waiting_sessions.insert(request_id, session); - } - - Ok(()) + if !accepts_both(&headers) { + return error( + StatusCode::NOT_ACCEPTABLE, + Value::Null, + -32600, + "Accept must include application/json and text/event-stream", + ); } - - fn activate_pending_calls( - &mut self, - channel_tx: &mut mpsc::UnboundedSender, - ) -> Result<(), agent_client_protocol::Error> { - loop { - let Some(call) = self.pending_calls.front() else { - return Ok(()); - }; - if call.session.outgoing_tx.is_closed() { - self.pending_calls.pop_front(); - continue; - } - if self.call_conflicts_with_active(&call.request_ids) { - return Ok(()); - } - - let call = self - .pending_calls - .pop_front() - .expect("pending call was checked above"); - self.activate_call(call, channel_tx)?; - } + if header_value(&headers, header::CONTENT_TYPE.as_str()) + .is_none_or(|value| !value.eq_ignore_ascii_case("application/json")) + { + return error( + StatusCode::UNSUPPORTED_MEDIA_TYPE, + Value::Null, + -32600, + "Expected application/json", + ); } - - fn call_conflicts_with_active(&self, request_ids: &[RequestId]) -> bool { - let unidentified_batch_is_active = self - .waiting_batch_sessions - .iter() - .any(|waiting| waiting.request_ids.is_empty()); - - if request_ids.is_empty() { - // A response-bearing batch without a request ID (for example, a - // notification plus an invalid scalar) receives a grouped error - // response whose only ID is null. Keep those responses ordered, - // including with explicit null-ID calls, because the wire response - // does not otherwise carry enough provenance to distinguish them. - return unidentified_batch_is_active || self.request_id_is_active(&RequestId::Null); + let body: Value = match serde_json::from_slice(&body) { + Ok(body) => body, + Err(_) => return error(StatusCode::BAD_REQUEST, Value::Null, -32700, "Parse error"), + }; + let id = body + .get("id") + .filter(|id| valid_request_id(id)) + .cloned() + .unwrap_or(Value::Null); + let Some(object) = body.as_object() else { + return error( + StatusCode::BAD_REQUEST, + id, + -32600, + "Expected one JSON-RPC request", + ); + }; + if object.get("jsonrpc").and_then(Value::as_str) != Some("2.0") + || object.contains_key("result") + || object.contains_key("error") + || object.get("id").is_none_or(|id| !valid_request_id(id)) + { + return error( + StatusCode::BAD_REQUEST, + id, + -32600, + "Expected one JSON-RPC request; batches, notifications and client responses are unsupported", + ); + } + let Some(method) = object.get("method").and_then(Value::as_str) else { + return error(StatusCode::BAD_REQUEST, id, -32600, "Missing method"); + }; + let params = match object.get("params") { + None => None, + Some(Value::Object(params)) => Some(params.clone()), + _ => { + return error( + StatusCode::BAD_REQUEST, + id, + -32602, + "Parameters must be an object", + ); } - - request_ids - .iter() - .any(|request_id| self.request_id_is_active(request_id)) - || unidentified_batch_is_active && request_ids.contains(&RequestId::Null) + }; + let metadata_version = params + .as_ref() + .and_then(|params| params.get("_meta")) + .and_then(Value::as_object) + .and_then(|meta| meta.get("io.modelcontextprotocol/protocolVersion")) + .and_then(Value::as_str); + let version = header_value(&headers, "mcp-protocol-version"); + if version.is_none() || version != metadata_version { + return error( + StatusCode::BAD_REQUEST, + id, + -32020, + "HeaderMismatch: MCP-Protocol-Version does not match params._meta", + ); } - - fn request_id_is_active(&self, request_id: &RequestId) -> bool { - self.waiting_sessions.contains_key(request_id) - || self - .waiting_batch_sessions - .iter() - .any(|waiting| waiting.request_ids.contains(request_id)) + if version != Some(VERSION) { + return ( + StatusCode::BAD_REQUEST, + Json(serde_json::json!({ + "jsonrpc":"2.0","id":id, + "error":{"code":-32022,"message":"Unsupported protocol version", + "data":{"supported":[VERSION],"requested":version}} + })), + ) + .into_response(); } - - fn drain_jsonrpc_messages(&mut self) { - while let Some(message) = self.message_deque.pop_front() { - if let Some(message) = self.try_dispatch_jsonrpc_message(message) { - self.message_deque.push_front(message); - break; - } - } + if header_value(&headers, "mcp-method") != Some(method) { + return error( + StatusCode::BAD_REQUEST, + id, + -32020, + "HeaderMismatch: Mcp-Method does not match method", + ); } - - fn try_dispatch_jsonrpc_message( - &mut self, - mut message: TransportFrame, - ) -> Option { - if matches!(message, TransportFrame::Malformed { .. }) { - // Malformed frames emitted by a relay are wire data, not protocol - // responses, so they are delivered through a general stream. - } else if matches!(message, TransportFrame::Batch(_)) { - let response_ids: Vec<_> = match &message { - TransportFrame::Batch(batch) => batch - .entries() - .filter_map(|entry| match entry { - TransportBatchEntry::Message(message) => message.response_id(), - TransportBatchEntry::Malformed { .. } => None, - }) - .collect(), - _ => unreachable!(), - }; - let correlated = self.waiting_batch_sessions.iter().position(|waiting| { - !waiting.request_ids.is_empty() - && waiting - .request_ids - .iter() - .any(|id| response_ids.contains(&id)) - }); - let fallback = response_ids.contains(&&RequestId::Null).then(|| { - self.waiting_batch_sessions - .iter() - .position(|waiting| waiting.request_ids.is_empty()) - }); - let fallback = fallback.flatten(); - if let Some(index) = correlated.or(fallback) { - let session = self.waiting_batch_sessions.remove(index).session; - // This response belongs to that HTTP POST even if its SSE - // receiver has gone away. Never let it fall through to a - // later request that reuses the same JSON-RPC ID. - drop(session.outgoing_tx.unbounded_send(message)); - return None; - } - } - - let message_id = match &message { - TransportFrame::Single(message) => message.response_id().cloned(), - TransportFrame::Malformed { .. } => None, - TransportFrame::Batch(batch) => batch.entries().find_map(|entry| match entry { - TransportBatchEntry::Message(message) => message.response_id().cloned(), - TransportBatchEntry::Malformed { .. } => None, - }), + if matches!(method, "tools/call" | "prompts/get" | "resources/read") { + let Some(name) = params + .as_ref() + .and_then(|params| mirrored_name(method, params)) + else { + return error( + StatusCode::BAD_REQUEST, + id, + -32602, + "Missing params.name or params.uri", + ); }; - - if let Some(ref message_id) = message_id - && let Some(session) = self.waiting_sessions.remove(message_id) - { - // This response belongs to that HTTP POST even if its SSE - // receiver has gone away. Never let it fall through to a later - // request that reuses the same JSON-RPC ID. - drop(session.outgoing_tx.unbounded_send(message)); - return None; - } - - self.purge_closed_sessions(); - let all_sessions = self - .general_sessions - .iter_mut() - .chain(self.waiting_sessions.values_mut()) - .chain( - self.waiting_batch_sessions - .iter_mut() - .map(|waiting| &mut waiting.session), - ) - .chain( - self.pending_calls - .iter_mut() - .map(|waiting| &mut waiting.session), + if !matches_mirror(header_value(&headers, "mcp-name"), name) { + return error( + StatusCode::BAD_REQUEST, + id, + -32020, + "HeaderMismatch: Mcp-Name does not match request", ); - for session in all_sessions { - match session.outgoing_tx.unbounded_send(message) { - Ok(()) => return None, - Err(m) => { - assert!(m.is_disconnected()); - message = m.into_inner(); - } - } } - - Some(message) } - - fn purge_closed_sessions(&mut self) { - self.general_sessions - .retain(|session| !session.outgoing_tx.is_closed()); - self.pending_calls - .retain(|call| !call.session.outgoing_tx.is_closed()); - - // Calls already forwarded to the JSON-RPC peer stay registered until - // their response arrives. Otherwise a late response could be routed - // to a newer HTTP POST that reused the same request ID. + // Tool schemas with x-mcp-header annotations are not tracked in this adapter. + // Fail closed on supplied mirrored parameter headers; support for annotations + // requires a request-scoped schema lookup and validation before forwarding. + if headers + .keys() + .any(|key| key.as_str().starts_with("mcp-param-")) + { + return error( + StatusCode::BAD_REQUEST, + id, + -32020, + "HeaderMismatch: Mcp-Param headers are not supported by this adapter", + ); } -} - -struct PendingCall { - frame: TransportFrame, - request_ids: Vec, - session: RegisteredSession, -} - -struct WaitingBatchSession { - request_ids: Vec, - session: RegisteredSession, -} - -struct RegisteredSession { - #[allow(dead_code)] - id: uuid::Uuid, - outgoing_tx: mpsc::UnboundedSender, -} - -impl RegisteredSession { - fn new(outgoing_tx: mpsc::UnboundedSender) -> Self { - Self { - id: uuid::Uuid::new_v4(), - outgoing_tx, - } + if method == "initialize" || method.starts_with("notifications/") { + return error(StatusCode::NOT_FOUND, id, -32601, "Method not found"); } -} - -/// Accept a POST request carrying a JSON-RPC frame from an MCP client. -/// For response-bearing calls and batches, we return an SSE stream. For -/// notification/response-only frames, we return 202 Accepted. -async fn handle_post( - State(state): State>, - body: String, -) -> Result { - let http_request_id = uuid::Uuid::new_v4(); - let frame = TransportFrame::parse_json(&body); - - match frame { - TransportFrame::Single(message) => match message { - RawJsonRpcMessage::Request(request) => { - let (tx, rx) = mpsc::unbounded(); - state - .registration_tx - .unbounded_send(HttpMessage::Request { - http_request_id, - request, - response_tx: tx, - }) - .map_err(agent_client_protocol::util::internal_error)?; - - Ok(sse_response(rx)) - } - RawJsonRpcMessage::Notification(request) => { - state - .registration_tx - .unbounded_send(HttpMessage::Notification { - http_request_id, - request, - }) - .map_err(agent_client_protocol::util::internal_error)?; - Ok(StatusCode::ACCEPTED.into_response()) - } - RawJsonRpcMessage::Response(response) => { - state - .registration_tx - .unbounded_send(HttpMessage::Response { - http_request_id, - response, - }) - .map_err(agent_client_protocol::util::internal_error)?; - Ok(StatusCode::ACCEPTED.into_response()) - } - }, - TransportFrame::Malformed { raw, error } => { - if raw - .parse::() - .is_ok_and(|value| is_response_only_shape(&value)) - { - return Ok(StatusCode::ACCEPTED.into_response()); - } - Ok(immediate_sse_response(TransportFrame::Single( - RawJsonRpcMessage::response(RequestId::Null, Err(error)), - ))) - } - TransportFrame::Batch(batch) => { - if batch - .entries() - .all(|entry| matches!(entry, TransportBatchEntry::Malformed { .. })) - { - let responses = agent_client_protocol::TransportBatch::from_messages( - batch.entries().filter_map(|entry| { - let TransportBatchEntry::Malformed { raw, error } = entry else { - unreachable!("all batch entries were checked as malformed") - }; - (!is_response_only_shape(raw)).then(|| { - RawJsonRpcMessage::response(RequestId::Null, Err(error.clone())) - }) - }), - ); - let Some(responses) = responses else { - return Ok(StatusCode::ACCEPTED.into_response()); - }; - return Ok(immediate_sse_response(TransportFrame::Batch(responses))); - } - - let mut request_ids = Vec::new(); - let mut expects_response = false; - for entry in batch.entries() { - match entry { - TransportBatchEntry::Message(RawJsonRpcMessage::Request(request)) => { - request_ids.push(request.id.clone()); - expects_response = true; - } - TransportBatchEntry::Malformed { raw, .. } => { - expects_response |= !is_response_only_shape(raw); - } - TransportBatchEntry::Message( - RawJsonRpcMessage::Notification(_) | RawJsonRpcMessage::Response(_), - ) => {} - } - } - let frame = TransportFrame::Batch(batch); - if expects_response { - let (tx, rx) = mpsc::unbounded(); - state - .registration_tx - .unbounded_send(HttpMessage::Frame { - http_request_id, - frame, - request_ids, - response_tx: Some(tx), - }) - .map_err(agent_client_protocol::util::internal_error)?; - Ok(sse_response(rx)) - } else { - state - .registration_tx - .unbounded_send(HttpMessage::Frame { - http_request_id, - frame, - request_ids, - response_tx: None, - }) - .map_err(agent_client_protocol::util::internal_error)?; - Ok(StatusCode::ACCEPTED.into_response()) - } - } + let (notification_tx, mut response_rx) = tokio_mpsc::channel(super::MAX_QUEUED_NOTIFICATIONS); + let response_tx = super::StreamSender { + tx: notification_tx, + used: Arc::new(std::sync::atomic::AtomicUsize::new(0)), + }; + let (terminal_tx, mut terminal_rx) = oneshot::channel(); + let message = BridgeMessage::Request { + server_id: state.server_id.clone(), + request_id: uuid::Uuid::new_v4().to_string(), + http_id: id, + method: method.into(), + params, + response_tx, + terminal_tx, + }; + let mut tx = state.tx.clone(); + if tx.send(message).await.is_err() { + return error( + StatusCode::SERVICE_UNAVAILABLE, + Value::Null, + -32603, + "ACP bridge unavailable", + ); } -} - -fn is_response_only_shape(value: &serde_json::Value) -> bool { - value.as_object().is_some_and(|object| { - !object.contains_key("method") - && (object.contains_key("result") || object.contains_key("error")) - }) -} - -/// Accept a GET request from an MCP client. -/// Opens an SSE stream for server-initiated messages. -async fn handle_get( - State(state): State>, -) -> Result>>, HttpError> { - let http_request_id = uuid::Uuid::new_v4(); - let (tx, mut rx) = mpsc::unbounded(); - state - .registration_tx - .unbounded_send(HttpMessage::Get { - http_request_id, - response_tx: tx, - }) - .map_err(agent_client_protocol::util::internal_error)?; - - let stream = async_stream::stream! { - while let Some(message) = rx.next().await { - yield sse_event(message); - } + let first = tokio::select! { + biased; + notification = response_rx.recv(), if !response_rx.is_closed() || !response_rx.is_empty() => + match notification { + Some(mut message) => Some(std::mem::take(&mut message.value)), + None => (&mut terminal_rx).await.ok(), + }, + terminal = &mut terminal_rx => terminal.ok(), }; - - Ok(Sse::new(stream)) -} - -fn sse_event(frame: TransportFrame) -> Result { - Ok(axum::response::sse::Event::default().data(frame.to_json()?)) -} - -fn sse_response(mut rx: mpsc::UnboundedReceiver) -> Response { + let Some(first) = first else { + return error( + StatusCode::SERVICE_UNAVAILABLE, + Value::Null, + -32603, + "ACP bridge closed", + ); + }; + if first.get("id").is_some() { + let status = if first.pointer("/error/code").and_then(Value::as_i64) == Some(-32601) { + StatusCode::NOT_FOUND + } else { + StatusCode::OK + }; + return (status, Json(first)).into_response(); + } let stream = async_stream::stream! { - while let Some(message) = rx.next().await { - yield sse_event(message); + yield Ok::<_, Infallible>(Event::default().data(first.to_string())); + loop { + // Drain already-queued notifications before a successful final response. + // Overflow is delivered through the independent terminal path. + let message = tokio::select! { + biased; + notification = response_rx.recv(), if !response_rx.is_closed() || !response_rx.is_empty() => + match notification { + Some(mut message) => Some(std::mem::take(&mut message.value)), + None => (&mut terminal_rx).await.ok(), + }, + terminal = &mut terminal_rx => terminal.ok(), + }; + let Some(message) = message else { break }; + let final_response = message.get("id").is_some(); + yield Ok::<_, Infallible>(Event::default().data(message.to_string())); + if final_response { break } } }; - Sse::new(stream).into_response() -} - -fn immediate_sse_response(frame: TransportFrame) -> Response { - Sse::new(futures::stream::once(async move { sse_event(frame) })).into_response() + let mut response = Sse::new(stream) + .keep_alive(KeepAlive::default()) + .into_response(); + response + .headers_mut() + .insert("x-accel-buffering", "no".parse().expect("static header")); + response } #[cfg(test)] mod tests { use super::*; - - async fn single_sse_payload(response: Response) -> serde_json::Value { - let body = axum::body::to_bytes(response.into_body(), 64 * 1024) - .await - .expect("SSE response body"); - let body = std::str::from_utf8(&body).expect("UTF-8 SSE response"); - let payload = body - .lines() - .find_map(|line| line.strip_prefix("data:").map(str::trim_start)) - .expect("one SSE data event"); - serde_json::from_str(payload).expect("JSON-RPC SSE payload") - } - - async fn single_sse_message(response: Response) -> RawJsonRpcMessage { - serde_json::from_value(single_sse_payload(response).await) - .expect("single JSON-RPC SSE message") - } + use tokio::io::{AsyncReadExt, AsyncWriteExt}; #[test] - fn malformed_post_cannot_steal_a_valid_null_id_response() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - - let valid_http_response = handle_post( - State(state.clone()), - serde_json::json!({ - "jsonrpc": "2.0", - "id": null, - "method": "example", - "params": {} - }) - .to_string(), - ) - .await - .expect("valid null-ID POST"); - let malformed_http_response = handle_post(State(state), "{not json".to_owned()) - .await - .expect("malformed POST receives a JSON-RPC error"); - - let valid_request = registration_rx - .next() - .await - .expect("valid request is forwarded"); - assert!( - registration_rx.try_recv().is_err(), - "malformed input must be answered by its own HTTP request" - ); - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - server - .handle_http_message(valid_request, &mut channel_tx) - .expect("forward valid request"); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(_))) - )); - - let valid_response = TransportFrame::Single(RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "valid" })), - )); - assert!( - server - .try_dispatch_jsonrpc_message(valid_response) - .is_none() - ); - - assert!(matches!( - single_sse_message(valid_http_response).await, - RawJsonRpcMessage::Response(RpcResponse::Result { - id: RequestId::Null, - result, - .. - }) if result == serde_json::json!({ "source": "valid" }) - )); - assert!(matches!( - single_sse_message(malformed_http_response).await, - RawJsonRpcMessage::Response(RpcResponse::Error { - id: RequestId::Null, - error, - .. - }) if error.code == agent_client_protocol::ErrorCode::ParseError - )); - }); + fn accepts_only_both_media_types() { + let mut headers = HeaderMap::new(); + headers.insert( + "accept", + "application/json, text/event-stream".parse().unwrap(), + ); + assert!(accepts_both(&headers)); + headers.insert("accept", "application/json".parse().unwrap()); + assert!(!accepts_both(&headers)); } #[test] - fn malformed_response_shaped_posts_are_ignored_without_registration() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - let malformed_response = serde_json::json!({ - "jsonrpc": "2.0", - "id": 1, - "result": null, - "error": { "code": -32603, "message": "Internal error" } - }); - - let single = handle_post(State(state.clone()), malformed_response.to_string()) - .await - .expect("malformed response-shaped POST"); - assert_eq!(single.status(), StatusCode::ACCEPTED); - - let batch = handle_post( - State(state), - serde_json::Value::Array(vec![malformed_response]).to_string(), - ) - .await - .expect("malformed response-only batch POST"); - assert_eq!(batch.status(), StatusCode::ACCEPTED); - assert!( - registration_rx.try_recv().is_err(), - "ignored responses must not be forwarded or register HTTP waiters" - ); - }); + fn mirrored_names_decode_canonical_base64() { + assert!(matches_mirror( + Some("=?base64?SGVsbG8sIOS4lueVjA==?="), + "Hello, 世界" + )); + assert!(matches_mirror(Some("simple"), "simple")); + assert!(!matches_mirror(Some("=?base64?SGVsbG8=?="), "different")); + assert!(!matches_mirror(Some("=?base64?SGVsbG8==?="), "Hello")); + assert!(!matches_mirror( + Some("=?base64?literal?="), + "=?base64?literal?=" + )); } #[test] - fn malformed_response_sibling_does_not_hide_invalid_batch_value() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - let response = handle_post( - State(state), - serde_json::json!([ - 17, - { - "jsonrpc": "2.0", - "id": 1, - "result": null, - "error": { "code": -32603, "message": "Internal error" } - } - ]) - .to_string(), - ) - .await - .expect("mixed malformed batch POST"); - - let payload = single_sse_payload(response).await; - let entries = payload.as_array().expect("batch response array"); - assert_eq!(entries.len(), 1); - assert_eq!(entries[0]["id"], serde_json::Value::Null); - assert_eq!( - entries[0]["error"]["code"], - i32::from(agent_client_protocol::ErrorCode::InvalidRequest) - ); - assert!( - registration_rx.try_recv().is_err(), - "an all-malformed batch is answered by its originating POST" - ); + fn response_preserves_mrtr_and_opaque_request_state() { + let result = serde_json::json!({ + "inputRequests": [{"method":"elicitation/create","params":{"message":"answer"}}], + "requestState": {"opaque": [1, 2, 3]}, + "_meta": {"trace":"preserve", "io.modelcontextprotocol/subscriptionId":"unrelated"}, + "subscriptionId": "internal-id" }); + let response = rpc_result(serde_json::json!(42), "internal-id", result.clone()); + assert_eq!(response["result"], result); + assert_eq!(response["id"], 42); + let mapped = rpc_result( + serde_json::json!("external"), + "internal-id", + serde_json::json!({"subscriptionId":"internal-id","requestState":"unchanged", + "_meta":{"trace":"preserve", "io.modelcontextprotocol/subscriptionId":"internal-id", + "progressToken":"internal-id"}}), + ); + assert_eq!(mapped["result"]["subscriptionId"], "internal-id"); + assert_eq!( + mapped["result"]["_meta"]["io.modelcontextprotocol/subscriptionId"], + "external" + ); + assert_eq!(mapped["result"]["_meta"]["progressToken"], "internal-id"); + assert_eq!(mapped["result"]["requestState"], "unchanged"); } #[test] - fn concurrent_null_id_posts_are_serialized_to_preserve_response_provenance() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - - let first_http_response = handle_post( - State(state.clone()), - serde_json::json!({ - "jsonrpc": "2.0", - "id": null, - "method": "example/first", - "params": {} - }) - .to_string(), - ) - .await - .expect("first null-ID POST"); - let second_http_response = handle_post( - State(state), - serde_json::json!({ - "jsonrpc": "2.0", - "id": null, - "method": "example/second", - "params": {} - }) - .to_string(), - ) - .await - .expect("second null-ID POST"); - - let first_registration = registration_rx.next().await.unwrap(); - let second_registration = registration_rx.next().await.unwrap(); - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - - server - .handle_http_message(first_registration, &mut channel_tx) - .unwrap(); - server - .handle_http_message(second_registration, &mut channel_tx) - .unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) - if request.method.as_ref() == "example/first" - && request.id == RequestId::Null - )); - assert!( - channel_rx.try_recv().is_err(), - "an overlapping null-ID request must wait for the first response" - ); - - assert!( - server - .try_dispatch_jsonrpc_message(TransportFrame::Single( - RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "first" })), - ), - )) - .is_none() - ); - server.activate_pending_calls(&mut channel_tx).unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) - if request.method.as_ref() == "example/second" - && request.id == RequestId::Null - )); - - assert!( - server - .try_dispatch_jsonrpc_message(TransportFrame::Single( - RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "second" })), - ), - )) - .is_none() - ); - - assert!(matches!( - single_sse_message(first_http_response).await, - RawJsonRpcMessage::Response(RpcResponse::Result { - id: RequestId::Null, - result, - .. - }) if result == serde_json::json!({ "source": "first" }) - )); - assert!(matches!( - single_sse_message(second_http_response).await, - RawJsonRpcMessage::Response(RpcResponse::Result { - id: RequestId::Null, - result, - .. - }) if result == serde_json::json!({ "source": "second" }) - )); - }); + fn concurrent_logical_ids_with_same_external_id_stay_request_scoped() { + let mut first = + serde_json::json!({"_meta":{"io.modelcontextprotocol/subscriptionId":"one"}}); + let mut second = + serde_json::json!({"_meta":{"io.modelcontextprotocol/subscriptionId":"two"}}); + rewrite_subscription_id(&mut first, "one", &serde_json::json!(7)); + rewrite_subscription_id(&mut second, "two", &serde_json::json!(7)); + assert_eq!(first["_meta"]["io.modelcontextprotocol/subscriptionId"], 7); + assert_eq!(second["_meta"]["io.modelcontextprotocol/subscriptionId"], 7); + let mut mismatch = + serde_json::json!({"_meta":{"io.modelcontextprotocol/subscriptionId":"two"}}); + rewrite_subscription_id(&mut mismatch, "one", &serde_json::json!("7")); + assert_eq!( + mismatch["_meta"]["io.modelcontextprotocol/subscriptionId"], + "two" + ); } - #[test] - fn unidentified_batch_posts_are_serialized_to_preserve_response_provenance() { - futures::executor::block_on(async { - fn unidentified_batch(method: &str) -> TransportFrame { - TransportFrame::parse_json( - &serde_json::json!([ - { - "jsonrpc": "2.0", - "method": method, - "params": {} - }, - 17 - ]) - .to_string(), - ) - } - - fn grouped_response(source: &str) -> TransportFrame { - TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": source })), - ), - ]) - .expect("grouped response is non-empty"), - ) - } - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - let (first_tx, mut first_rx) = mpsc::unbounded(); - let (second_tx, mut second_rx) = mpsc::unbounded(); - - let first_frame = unidentified_batch("example/first"); - let second_frame = unidentified_batch("example/second"); - let expected_first_frame = first_frame.to_json().unwrap(); - let expected_second_frame = second_frame.to_json().unwrap(); - server - .handle_http_message( - HttpMessage::Frame { - http_request_id: uuid::Uuid::new_v4(), - frame: first_frame, - request_ids: Vec::new(), - response_tx: Some(first_tx), - }, - &mut channel_tx, - ) - .unwrap(); - server - .handle_http_message( - HttpMessage::Frame { - http_request_id: uuid::Uuid::new_v4(), - frame: second_frame, - request_ids: Vec::new(), - response_tx: Some(second_tx), - }, - &mut channel_tx, - ) - .unwrap(); - - assert_eq!( - channel_rx.next().await.unwrap().to_json().unwrap(), - expected_first_frame + #[tokio::test] + async fn rejects_legacy_methods_and_invalid_headers_over_real_http() { + async fn exchange( + address: std::net::SocketAddr, + method: &str, + headers: &str, + body: &str, + ) -> String { + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + let request = format!( + "{method} / HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n{headers}Content-Length: {}\r\n\r\n{body}", + body.len() ); - assert!( - channel_rx.try_recv().is_err(), - "a second unidentified batch must wait for the first response" - ); - - let callback = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::notification("callback".into(), serde_json::json!({})) - .unwrap(), - ]) - .expect("callback batch is non-empty"), - ); - assert!(server.try_dispatch_jsonrpc_message(callback).is_none()); - assert!(matches!( - first_rx.next().await, - Some(TransportFrame::Batch(_)) - )); - - let first_response = grouped_response("first"); - let expected_first_response = first_response.to_json().unwrap(); - assert!( - server - .try_dispatch_jsonrpc_message(first_response) - .is_none() - ); - assert_eq!( - first_rx.next().await.unwrap().to_json().unwrap(), - expected_first_response - ); - - server.activate_pending_calls(&mut channel_tx).unwrap(); - assert_eq!( - channel_rx.next().await.unwrap().to_json().unwrap(), - expected_second_frame - ); - - let second_response = grouped_response("second"); - let expected_second_response = second_response.to_json().unwrap(); - assert!( - server - .try_dispatch_jsonrpc_message(second_response) - .is_none() - ); - assert_eq!( - second_rx.next().await.unwrap().to_json().unwrap(), - expected_second_response - ); - }); - } - - #[test] - fn late_response_to_disconnected_post_cannot_reach_reused_id() { - futures::executor::block_on(async { - fn request(method: &str) -> RpcRequest { - let RawJsonRpcMessage::Request(request) = RawJsonRpcMessage::request( - method.to_owned(), - serde_json::json!({}), - RequestId::Null, - ) - .unwrap() else { - unreachable!("request constructor always returns a request") - }; - request - } - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - let (first_tx, first_rx) = mpsc::unbounded(); - let (second_tx, mut second_rx) = mpsc::unbounded(); - - server - .handle_http_message( - HttpMessage::Request { - http_request_id: uuid::Uuid::new_v4(), - request: request("example/first"), - response_tx: first_tx, - }, - &mut channel_tx, - ) - .unwrap(); - server - .handle_http_message( - HttpMessage::Request { - http_request_id: uuid::Uuid::new_v4(), - request: request("example/second"), - response_tx: second_tx, - }, - &mut channel_tx, - ) - .unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) - if request.method.as_ref() == "example/first" - )); - drop(first_rx); - - let callback = TransportFrame::Single( - RawJsonRpcMessage::notification("callback".into(), serde_json::json!({})).unwrap(), - ); - assert!(server.try_dispatch_jsonrpc_message(callback).is_none()); - assert!(matches!( - second_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Notification(_))) - )); - - let first_response = TransportFrame::Single(RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "first" })), - )); - assert!( - server - .try_dispatch_jsonrpc_message(first_response) - .is_none() - ); - server.activate_pending_calls(&mut channel_tx).unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) - if request.method.as_ref() == "example/second" - )); - - let second_response = TransportFrame::Single(RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "second" })), - )); - assert!( - server - .try_dispatch_jsonrpc_message(second_response) - .is_none() - ); - assert!(matches!( - second_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Response( - RpcResponse::Result { result, .. } - ))) if result == serde_json::json!({ "source": "second" }) - )); - }); - } - - #[test] - fn forwards_batch_and_routes_grouped_response_without_flattening() { - futures::executor::block_on(async { - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - let (response_tx, mut response_rx) = mpsc::unbounded(); - let incoming = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([RawJsonRpcMessage::request( - "example".into(), - serde_json::json!({}), - RequestId::Number(7), - ) - .unwrap()]) - .unwrap(), - ); - - server - .handle_http_message( - HttpMessage::Frame { - http_request_id: uuid::Uuid::new_v4(), - frame: incoming, - request_ids: vec![RequestId::Number(7)], - response_tx: Some(response_tx), - }, - &mut channel_tx, - ) - .unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Batch(_)) - )); - - let callback = TransportFrame::Single( - RawJsonRpcMessage::notification("callback".into(), serde_json::json!({})).unwrap(), - ); - assert!(server.try_dispatch_jsonrpc_message(callback).is_none()); - assert!(matches!( - response_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Notification(_))) - )); - - let frame = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::response( - RequestId::Number(7), - Ok(serde_json::json!({ "ok": true })), - ), - ]) - .unwrap(), - ); - let expected = frame.to_json().unwrap(); - - assert!(server.try_dispatch_jsonrpc_message(frame).is_none()); - let received = response_rx - .next() - .await - .expect("waiting HTTP request stays open"); - assert_eq!(received.to_json().unwrap(), expected); - assert!(matches!(received, TransportFrame::Batch(_))); - }); - } - - #[test] - fn batch_post_round_trips_as_one_grouped_sse_response() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - let incoming = serde_json::json!([ - { - "jsonrpc": "2.0", - "id": 7, - "method": "example/first", - "params": {} - }, - { - "jsonrpc": "2.0", - "id": 8, - "method": "example/second", - "params": {} - } - ]); - - let http_response = handle_post(State(state), incoming.to_string()) - .await - .expect("batch POST should open an SSE response"); - let registration = registration_rx - .next() - .await - .expect("batch POST should register with the bridge"); - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - server - .handle_http_message(registration, &mut channel_tx) - .expect("batch POST should be forwarded to the channel"); - let forwarded = channel_rx - .next() - .await - .expect("channel should receive the batch frame"); - assert!(matches!(&forwarded, TransportFrame::Batch(_))); - assert_eq!( - serde_json::from_str::(&forwarded.to_json().unwrap()).unwrap(), - incoming - ); - - let response = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::response( - RequestId::Number(7), - Ok(serde_json::json!({ "source": "first" })), - ), - RawJsonRpcMessage::response( - RequestId::Number(8), - Ok(serde_json::json!({ "source": "second" })), - ), - ]) - .expect("grouped response should be non-empty"), - ); - assert!(server.try_dispatch_jsonrpc_message(response).is_none()); - - let payload = single_sse_payload(http_response).await; - let entries = payload - .as_array() - .expect("SSE payload should remain one JSON-RPC array"); - assert_eq!(entries.len(), 2); - assert_eq!(entries[0]["id"], 7); - assert_eq!(entries[0]["result"]["source"], "first"); - assert_eq!(entries[1]["id"], 8); - assert_eq!(entries[1]["result"]["source"], "second"); - }); - } - - #[test] - fn malformed_batch_cannot_steal_a_valid_null_id_batch_response() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - - let valid_http_response = handle_post( - State(state.clone()), - serde_json::json!([{ - "jsonrpc": "2.0", - "id": null, - "method": "example", - "params": {} - }]) - .to_string(), - ) - .await - .expect("valid null-ID batch POST"); - let malformed_http_response = handle_post(State(state), "[17,false]".to_owned()) - .await - .expect("malformed batch receives its own JSON-RPC error array"); - - let valid_batch = registration_rx - .next() - .await - .expect("valid batch is forwarded"); - assert!( - registration_rx.try_recv().is_err(), - "malformed-only batch must not register a bridge waiter" - ); - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - server - .handle_http_message(valid_batch, &mut channel_tx) - .expect("forward valid null-ID batch"); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Batch(_)) - )); - - let valid_response = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "valid" })), - ), - ]) - .expect("valid response batch is non-empty"), - ); - assert!( - server - .try_dispatch_jsonrpc_message(valid_response) - .is_none() - ); - - let valid_payload = single_sse_payload(valid_http_response).await; - let valid_entries = valid_payload - .as_array() - .expect("valid response should remain a batch"); - assert_eq!(valid_entries.len(), 1); - assert_eq!(valid_entries[0]["id"], serde_json::Value::Null); - assert_eq!(valid_entries[0]["result"]["source"], "valid"); - - let malformed_payload = single_sse_payload(malformed_http_response).await; - let malformed_entries = malformed_payload - .as_array() - .expect("malformed response should be an error batch"); - assert_eq!(malformed_entries.len(), 2); - for entry in malformed_entries { - assert_eq!(entry["id"], serde_json::Value::Null); - assert_eq!( - entry["error"]["code"], - i32::from(agent_client_protocol::ErrorCode::InvalidRequest) - ); - } - }); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut response = String::new(); + stream.read_to_string(&mut response).await.unwrap(); + response + } + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let (tx, _rx) = mpsc::channel(8); + let task = tokio::spawn(run_http_listener( + listener, + "server".into(), + "secret".into(), + tx, + )); + let legacy = exchange(address, "GET", "", "").await; + assert!(legacy.starts_with("HTTP/1.1 405"), "{legacy}"); + let delete = exchange(address, "DELETE", "", "").await; + assert!(delete.starts_with("HTTP/1.1 405"), "{delete}"); + let invalid_origin = exchange(address, "POST", "Origin: http://evil.test\r\n", "{}").await; + assert!( + invalid_origin.starts_with("HTTP/1.1 403"), + "{invalid_origin}" + ); + let invalid_auth = exchange(address, "POST", "", "{}").await; + assert!(invalid_auth.starts_with("HTTP/1.1 401"), "{invalid_auth}"); + let body = serde_json::json!({"jsonrpc":"2.0","id":1,"method":"tools/list", + "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":VERSION, + "io.modelcontextprotocol/clientCapabilities":{}}}}) + .to_string(); + let headers = "Authorization: Bearer secret\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: wrong/method\r\n"; + let mismatch = exchange(address, "POST", headers, &body).await; + assert!(mismatch.starts_with("HTTP/1.1 400"), "{mismatch}"); + assert!(mismatch.contains("-32020"), "{mismatch}"); + let batch = exchange(address, "POST", headers, "[]").await; + assert!(batch.starts_with("HTTP/1.1 400"), "{batch}"); + let headers = headers.replace("wrong/method", "tools/list"); + let fractional_id = body.replace("\"id\":1", "\"id\":1.5"); + let fractional = tokio::time::timeout( + std::time::Duration::from_secs(3), + exchange(address, "POST", &headers, &fractional_id), + ) + .await + .expect("an invalid request ID must be rejected before forwarding"); + assert!(fractional.starts_with("HTTP/1.1 400"), "{fractional}"); + assert!(fractional.contains("-32600"), "{fractional}"); + let error: Value = + serde_json::from_str(fractional.split("\r\n\r\n").nth(1).unwrap()).unwrap(); + assert!(error.get("id").is_none()); + task.abort(); } } diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs index feeb6624..040f4f6d 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs @@ -1,121 +1,114 @@ -//! MCP-over-ACP compatibility proxy. +//! Request-scoped MCP 2026-07-28 Streamable HTTP adapter for native ACP MCP servers. //! -//! This proxy adapts schema-native `McpServer::Acp` declarations for agents that do not -//! support the ACP MCP transport. It replaces those declarations with loopback HTTP bridges and -//! relays `mcp/connect`, `mcp/message`, and `mcp/disconnect` over ACP. -//! -//! Stable protocol v1 is supported by default. Enable the crate's -//! `unstable_protocol_v2` feature to use the same proxy in a draft-v2 conductor -//! chain. -//! -//! # Usage -//! -//! ```rust,ignore -//! use agent_client_protocol_polyfill::mcp_over_acp::McpOverAcpPolyfill; -//! -//! let conductor = ConductorImpl::new_agent( -//! "conductor", -//! ProxiesAndAgent::new(my_agent).proxy(McpOverAcpPolyfill::http()), -//! ); -//! ``` +//! Native-capable successors receive the original declarations and messages unchanged. +//! HTTP-only successors receive loopback endpoints; no MCP connection or session is created. -mod actor; pub(crate) mod http; mod protocol; -use std::collections::HashMap; +use std::{ + collections::{HashMap, HashSet}, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, +}; use agent_client_protocol::{ Agent, Client, Conductor, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - Proxy, Responder, UntypedMessage, is_cancel_request_notification, util::MatchDispatchFrom, + Proxy, UntypedMessage, util::MatchDispatchFrom, +}; +use futures::{ + SinkExt, StreamExt, + channel::{mpsc, oneshot}, }; -use futures::{SinkExt, channel::mpsc, channel::oneshot}; use serde_json::Value; -use tokio::net::TcpListener; -use tracing::{debug, info, warn}; +use tokio::{net::TcpListener, sync::mpsc as tokio_mpsc}; +use tracing::{debug, warn}; + +use self::protocol::{DownstreamMcpMode, NativeMcpNotification, NativeServer, PolyfillProtocol}; + +// Conservative per-bridge limits. Notifications are bounded per HTTP POST by +// both message count and serialized bytes; terminal responses bypass the queue. +const MAX_ACTIVE_REQUESTS: usize = 64; +const MAX_LISTENERS: usize = 32; +const MAX_QUEUED_NOTIFICATIONS: usize = 16; +const MAX_QUEUED_BYTES: usize = 256 * 1024; + +struct QueuedNotification { + value: Value, + bytes: usize, + used: Arc, +} -use self::actor::BridgeConnectionActor; -use self::protocol::{ - DownstreamMcpMode, NativeMcpMessage, NativeServer, PolyfillProtocol, native_params_into_value, -}; +impl Drop for QueuedNotification { + fn drop(&mut self) { + self.used.fetch_sub(self.bytes, Ordering::Relaxed); + } +} -/// Internal messages for the polyfill's bridge management. -#[derive(Debug)] -pub(crate) enum BridgeMessage { - /// Record the selected ACP schema and which MCP transport the successor can consume. +#[derive(Clone)] +struct StreamSender { + tx: tokio_mpsc::Sender, + used: Arc, +} + +impl StreamSender { + fn send(&self, value: Value) -> Result<(), ()> { + let bytes = serde_json::to_vec(&value).map_err(|_| ())?.len(); + let reserved = self + .used + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |used| { + used.checked_add(bytes) + .filter(|total| *total <= MAX_QUEUED_BYTES) + }); + if reserved.is_err() { + return Err(()); + } + self.tx + .try_send(QueuedNotification { + value, + bytes, + used: self.used.clone(), + }) + .map_err(|_| ()) + } + + async fn closed(&self) { + self.tx.closed().await; + } +} + +enum BridgeMessage { SetProtocol { protocol: PolyfillProtocol, downstream_mode: DownstreamMcpMode, }, - - /// Transform the MCP declarations for one session setup request. TransformServers { servers: Vec, response_tx: oneshot::Sender, agent_client_protocol::Error>>, }, - - /// A new TCP connection was accepted and needs a native MCP connection ID. - ConnectionReceived { + Request { server_id: String, - actor: BridgeConnectionActor, - connection: BridgeConnection, - }, - - /// A native MCP connection ID was received; spawn the actor and store its sender. - ConnectionEstablished { - server_id: String, - connection_id: String, - actor: BridgeConnectionActor, - connection: BridgeConnection, - }, - - /// Opening a native MCP connection failed. - ConnectionFailed { server_id: String }, - - /// An MCP message from the local agent that must be sent over ACP. - ClientToServer { - connection_id: String, - message: Dispatch, + request_id: String, + http_id: Value, + method: String, + params: Option>, + response_tx: StreamSender, + terminal_tx: tokio::sync::oneshot::Sender, }, - - /// An MCP server request received over ACP for the local agent's MCP client. - ServerToClientRequest { - request: NativeMcpMessage, - responder: Responder, + Notification(NativeMcpNotification), + Finished { + request_id: String, + result: Option>, }, - - /// An MCP server notification received over ACP for the local agent's MCP client. - ServerToClientNotification { notification: NativeMcpMessage }, - - /// The local MCP bridge disconnected. - Disconnected { connection_id: String }, -} - -/// Connection handle for sending messages to an MCP client via a bridge. -#[derive(Clone, Debug)] -pub(crate) struct BridgeConnection { - to_mcp_client_tx: mpsc::Sender, } -impl BridgeConnection { - pub fn new(to_mcp_client_tx: mpsc::Sender) -> Self { - Self { to_mcp_client_tx } - } - - fn try_send(&mut self, message: Dispatch) -> Option> { - self.to_mcp_client_tx - .try_send(message) - .err() - .map(|error| Box::new(error.into_inner())) - } -} - -/// Adapts schema-native MCP-over-ACP declarations for agents that support HTTP MCP. +/// Adapts native MCP-over-ACP servers to loopback Streamable HTTP for HTTP-only agents. #[derive(Debug, Default)] pub struct McpOverAcpPolyfill; impl McpOverAcpPolyfill { - /// Create a polyfill that exposes each ACP MCP server through loopback HTTP. #[must_use] pub fn http() -> Self { Self @@ -136,7 +129,6 @@ impl ConnectTo for McpOverAcpPolyfill { .connect_to(client) .await } - #[cfg(not(feature = "unstable_protocol_v2"))] { McpOverAcpProxy(PolyfillProtocol::V1) @@ -160,14 +152,13 @@ impl ConnectTo for McpOverAcpProxy { bridge_rx, protocol: None, downstream_mode: DownstreamMcpMode::Unknown, - listeners: BridgeListeners::default(), - bridge_connections: HashMap::new(), + listeners: HashMap::new(), + active: HashMap::new(), }; let handler = PolyfillHandler { protocol: None, bridge_tx, }; - match self.0 { PolyfillProtocol::V1 => { Proxy @@ -224,11 +215,81 @@ impl PolyfillHandler { cx: &ConnectionTo, ) -> Result, agent_client_protocol::Error> { match message { - Dispatch::Request(request, responder) => { - self.handle_client_request(request, responder, cx).await + Dispatch::Request(mut request, responder) => { + if request.method() == agent_client_protocol::schema::METHOD_INITIALIZE_PROXY { + if self.protocol.is_some() { + return Err(agent_client_protocol::Error::invalid_request() + .data("MCP-over-ACP polyfill was already initialized")); + } + let protocol = PolyfillProtocol::from_initialize_request(&request)?; + self.protocol = Some(protocol); + request.method = "initialize".into(); + let sent = cx + .send_request_to(Agent, request) + .forward_cancellation_from(responder.cancellation()); + let mut bridge_tx = self.bridge_tx.clone(); + sent.on_receiving_result(async move |result| { + let result = match result { + Ok(mut response) => { + let mode = protocol.transform_initialize_response(&mut response)?; + bridge_tx + .send(BridgeMessage::SetProtocol { + protocol, + downstream_mode: mode, + }) + .await + .map_err(agent_client_protocol::Error::into_internal_error)?; + Ok(response) + } + Err(error) => Err(error), + }; + responder.respond_with_result(result) + })?; + return Ok(Handled::Yes); + } + let Some(protocol) = self.protocol else { + return Ok(Handled::No { + message: Dispatch::Request(request, responder), + retry: false, + }); + }; + if protocol.is_session_setup_method(request.method()) { + protocol.validate_session_setup_request(&request)?; + transform_session_servers(&mut request, &mut self.bridge_tx).await?; + cx.send_request_to(Agent, request) + .forward_response_to(responder)?; + return Ok(Handled::Yes); + } + // Only agent-to-provider requests are valid; reverse RPC is never forwarded. + if request.method() == "mcp/message" { + responder + .respond_with_error(agent_client_protocol::Error::method_not_found())?; + return Ok(Handled::Yes); + } + Ok(Handled::No { + message: Dispatch::Request(request, responder), + retry: false, + }) } Dispatch::Notification(notification) => { - self.handle_client_notification(notification).await + if notification.method() == "mcp/message" { + let Some(protocol) = self.protocol else { + return Ok(Handled::No { + message: Dispatch::Notification(notification), + retry: false, + }); + }; + let notification = protocol.parse_notification(notification)?; + self.bridge_tx + .send(BridgeMessage::Notification(notification)) + .await + .map_err(agent_client_protocol::Error::into_internal_error)?; + return Ok(Handled::Yes); + } + Ok(Handled::No { + message: Dispatch::Notification(notification), + retry: false, + }) } message @ Dispatch::Response(_, _) => Ok(Handled::No { message, @@ -236,108 +297,6 @@ impl PolyfillHandler { }), } } - - async fn handle_client_request( - &mut self, - mut request: UntypedMessage, - responder: Responder, - cx: &ConnectionTo, - ) -> Result, agent_client_protocol::Error> { - if request.method() == agent_client_protocol::schema::METHOD_INITIALIZE_PROXY { - if self.protocol.is_some() { - return Err(agent_client_protocol::Error::invalid_request() - .data("MCP-over-ACP polyfill was already initialized")); - } - let protocol = PolyfillProtocol::from_initialize_request(&request)?; - self.protocol = Some(protocol); - request.method = "initialize".to_string(); - - let sent = cx.send_request_to(Agent, request); - let sent = sent.forward_cancellation_from(responder.cancellation()); - let mut bridge_tx = self.bridge_tx.clone(); - sent.on_receiving_result(async move |result| { - let result = match result { - Ok(response) => { - adapt_initialize_response(protocol, response, &mut bridge_tx).await - } - Err(error) => Err(error), - }; - responder.respond_with_result(result) - })?; - return Ok(Handled::Yes); - } - - let Some(protocol) = self.protocol else { - return Ok(Handled::No { - message: Dispatch::Request(request, responder), - retry: false, - }); - }; - - if protocol.is_session_setup_method(request.method()) { - protocol.validate_session_setup_request(&request)?; - transform_session_servers(&mut request, &mut self.bridge_tx).await?; - cx.send_request_to(Agent, request) - .forward_response_to(responder)?; - return Ok(Handled::Yes); - } - - if request.method() == "mcp/message" { - let request = protocol.parse_message_request(request)?; - self.bridge_tx - .send(BridgeMessage::ServerToClientRequest { request, responder }) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - return Ok(Handled::Yes); - } - - Ok(Handled::No { - message: Dispatch::Request(request, responder), - retry: false, - }) - } - - async fn handle_client_notification( - &mut self, - notification: UntypedMessage, - ) -> Result, agent_client_protocol::Error> { - let Some(protocol) = self.protocol else { - return Ok(Handled::No { - message: Dispatch::Notification(notification), - retry: false, - }); - }; - - if notification.method() == "mcp/message" { - let notification = protocol.parse_message_notification(notification)?; - self.bridge_tx - .send(BridgeMessage::ServerToClientNotification { notification }) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - return Ok(Handled::Yes); - } - - Ok(Handled::No { - message: Dispatch::Notification(notification), - retry: false, - }) - } -} - -async fn adapt_initialize_response( - protocol: PolyfillProtocol, - mut response: Value, - bridge_tx: &mut mpsc::Sender, -) -> Result { - let downstream_mode = protocol.transform_initialize_response(&mut response)?; - bridge_tx - .send(BridgeMessage::SetProtocol { - protocol, - downstream_mode, - }) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - Ok(response) } async fn transform_session_servers( @@ -352,7 +311,6 @@ async fn transform_session_servers( else { return Ok(()); }; - let (response_tx, response_rx) = oneshot::channel(); bridge_tx .send(BridgeMessage::TransformServers { @@ -367,14 +325,10 @@ async fn transform_session_servers( Ok(()) } -#[derive(Default, Debug)] -struct BridgeListeners { - listeners: HashMap, -} - -#[derive(Clone, Debug)] struct BridgeListener { tcp_port: u16, + // Runtime-only; never trace the listener or the rewritten declaration. + token: String, } impl BridgeListener { @@ -383,85 +337,21 @@ impl BridgeListener { protocol: PolyfillProtocol, server: NativeServer, ) -> Result { - server.http_declaration(protocol, format!("http://127.0.0.1:{}", self.tcp_port)) - } -} - -impl BridgeListeners { - async fn transform_servers( - &mut self, - connection: &ConnectionTo, - protocol: PolyfillProtocol, - servers: Vec, - bridge_tx: &mpsc::Sender, - ) -> Result, agent_client_protocol::Error> { - let mut transformed = Vec::with_capacity(servers.len()); - for server in servers { - transformed.push( - self.transform_server(connection, protocol, server, bridge_tx) - .await?, - ); - } - Ok(transformed) - } - - async fn transform_server( - &mut self, - connection: &ConnectionTo, - protocol: PolyfillProtocol, - server: Value, - bridge_tx: &mpsc::Sender, - ) -> Result { - let Some(native_server) = protocol.native_server(server.clone()) else { - return Ok(server); - }; - let server_id = native_server.server_id.clone(); - - info!( - server_name = %native_server.name, - server_id, - "detected native MCP-over-ACP server; creating compatibility bridge" - ); - - if let Some(listener) = self.listeners.get(&server_id) { - return listener.declaration(protocol, native_server); - } - - let tcp_listener = TcpListener::bind("127.0.0.1:0") - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - let tcp_port = tcp_listener - .local_addr() - .map_err(agent_client_protocol::Error::into_internal_error)? - .port(); - let listener = BridgeListener { tcp_port }; - - connection.spawn({ - let server_id = server_id.clone(); - let bridge_tx = bridge_tx.clone(); - async move { - info!( - server_id, - tcp_port, "accepting MCP compatibility connections" - ); - http::run_http_listener(tcp_listener, server_id, bridge_tx).await - } - })?; - - let declaration = listener.declaration(protocol, native_server)?; - self.listeners.insert(server_id, listener); - Ok(declaration) - } - - fn remove(&mut self, server_id: &str) { - self.listeners.remove(server_id); + server.http_declaration( + protocol, + format!("http://127.0.0.1:{}", self.tcp_port), + &self.token, + ) } } -#[derive(Debug)] -struct ActiveBridgeConnection { +struct ActiveRequest { server_id: String, - bridge: BridgeConnection, + http_id: Value, + method: String, + response_tx: StreamSender, + terminal_tx: tokio::sync::oneshot::Sender, + cancel_tx: tokio::sync::oneshot::Sender<()>, } struct BridgeRunner { @@ -469,8 +359,8 @@ struct BridgeRunner { bridge_rx: mpsc::Receiver, protocol: Option, downstream_mode: DownstreamMcpMode, - listeners: BridgeListeners, - bridge_connections: HashMap, + listeners: HashMap, + active: HashMap, } impl std::fmt::Debug for BridgeRunner { @@ -478,8 +368,8 @@ impl std::fmt::Debug for BridgeRunner { f.debug_struct("BridgeRunner") .field("protocol", &self.protocol) .field("downstream_mode", &self.downstream_mode) - .field("listeners", &self.listeners.listeners.len()) - .field("bridge_connections", &self.bridge_connections.len()) + .field("listeners", &self.listeners.len()) + .field("active", &self.active.len()) .finish_non_exhaustive() } } @@ -489,8 +379,6 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { mut self, connection: ConnectionTo, ) -> Result<(), agent_client_protocol::Error> { - use futures::StreamExt; - while let Some(message) = self.bridge_rx.next().await { match message { BridgeMessage::SetProtocol { @@ -500,281 +388,132 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { self.protocol = Some(protocol); self.downstream_mode = downstream_mode; } - BridgeMessage::TransformServers { servers, response_tx, } => { - let result = match (self.protocol, self.downstream_mode) { - (Some(_), DownstreamMcpMode::Native) => Ok(servers), - (Some(protocol), DownstreamMcpMode::HttpAdapter) => { - self.listeners - .transform_servers(&connection, protocol, servers, &self.bridge_tx) - .await - } - (Some(protocol), DownstreamMcpMode::Unavailable) => reject_native_servers( - protocol, - servers, - "the downstream agent supports neither native nor HTTP MCP transport", - ), - (Some(protocol), DownstreamMcpMode::Unknown) => reject_native_servers( - protocol, - servers, - "MCP transport capabilities are unavailable before initialize", - ), - (None, _) => Err(agent_client_protocol::Error::invalid_request() - .data("MCP transport capabilities are unavailable before initialize")), - }; + let result = self.transform_servers(&connection, servers).await; drop(response_tx.send(result)); } - - BridgeMessage::ConnectionReceived { + BridgeMessage::Request { server_id, - actor, - connection: bridge, + request_id, + http_id, + method, + params, + response_tx, + terminal_tx, } => { - let Some(protocol) = self.protocol else { - warn!( - server_id, - "cannot open MCP bridge before ACP initialization" - ); - self.listeners.remove(&server_id); + let Some(protocol) = self + .protocol + .filter(|_| self.downstream_mode == DownstreamMcpMode::HttpAdapter) + else { + drop(terminal_tx.send(http::rpc_error( + http_id, + -32603, + "MCP adapter unavailable", + ))); continue; }; - let request = protocol.connect_request(server_id.clone())?; - let mut bridge_tx = self.bridge_tx.clone(); - let scheduled = connection - .send_request_to(Client, request) - .on_receiving_result(async move |result| { - let message = match result { - Ok(response) => match protocol.connect_response_id(response) { - Ok(connection_id) => BridgeMessage::ConnectionEstablished { - server_id, - connection_id, - actor, - connection: bridge, - }, - Err(error) => { - warn!(?error, "invalid response to mcp/connect"); - BridgeMessage::ConnectionFailed { server_id } - } - }, - Err(error) => { - warn!(?error, "mcp/connect failed"); - BridgeMessage::ConnectionFailed { server_id } - } - }; - drop(bridge_tx.send(message).await); - Ok(()) - }); - if let Err(error) = scheduled { - warn!(?error, "could not schedule mcp/connect response handling"); - } - } - - BridgeMessage::ConnectionEstablished { - server_id, - connection_id, - actor, - connection: bridge, - } => { - self.bridge_connections.insert( - connection_id.clone(), - ActiveBridgeConnection { server_id, bridge }, - ); - connection.spawn(actor.run(connection_id))?; - } - - BridgeMessage::ConnectionFailed { server_id } => { - self.listeners.remove(&server_id); - } - - BridgeMessage::ClientToServer { - connection_id, - message, - } => { - let Some(protocol) = self.protocol else { - let rejection = match message { - Dispatch::Request(_, responder) => responder - .respond_with_internal_error( - "ACP protocol is unavailable before initialize", - ), - Dispatch::Notification(_) | Dispatch::Response(_, _) => Ok(()), - }; - if let Err(error) = rejection { - debug!(?error, "could not reject MCP request before initialize"); - } + if !self.listeners.contains_key(&server_id) { + drop(terminal_tx.send(http::rpc_error( + http_id, + -32602, + "Unknown MCP server", + ))); continue; - }; - - match message { - Dispatch::Request(message, responder) => { - match protocol.message_request(connection_id, message) { - Ok(request) => { - let pending = connection.send_request_to(Client, request); - if let Err(error) = pending.forward_response_to(responder) { - warn!( - ?error, - "could not forward local MCP request response" - ); - } - } - Err(error) => { - if let Err(send_error) = responder.respond_with_error(error) { - debug!( - ?send_error, - "could not reject malformed MCP request" - ); - } - } - } - } - Dispatch::Notification(message) => { - match local_mcp_notification(protocol, connection_id, message) { - Ok(Some(notification)) => { - if let Err(error) = - connection.send_notification_to(Client, notification) - { - warn!(?error, "could not forward local MCP notification"); - } - } - Ok(None) => { - debug!( - "not tunneling hop-scoped MCP cancellation through mcp/message" - ); - } - Err(error) => { - warn!(?error, "could not forward local MCP notification"); - } - } - } - Dispatch::Response(result, router) => { - if let Err(error) = router.route_with_result(result) { - debug!(?error, "could not route MCP client response"); - } - } } - } - - BridgeMessage::ServerToClientRequest { request, responder } => { - match self.downstream_mode { - DownstreamMcpMode::Native => { - let pending = connection.send_request_to(Agent, request.raw); - if let Err(error) = pending.forward_response_to(responder) { - debug!(?error, "could not forward native MCP request"); - } - } - DownstreamMcpMode::HttpAdapter => { - let connection_id = request.connection_id; - let Some(active) = self.bridge_connections.get_mut(&connection_id) - else { - respond_unknown_connection(responder, &connection_id); - continue; - }; - let message = UntypedMessage { - method: request.method, - params: native_params_into_value(request.params), - }; - if let Some(message) = active - .bridge - .try_send(Dispatch::Request(message, responder)) - { - let Dispatch::Request(_, responder) = *message else { - unreachable!("the failed bridge message was a request") - }; - if let Err(send_error) = responder.respond_with_internal_error( - "the local MCP client is unavailable or backpressured", - ) { - debug!( - ?send_error, - "could not reject unavailable MCP connection" - ); - } - } - } - DownstreamMcpMode::Unknown | DownstreamMcpMode::Unavailable => { - if let Err(error) = - responder.respond_with_error( - agent_client_protocol::Error::method_not_found(), - ) - { - debug!(?error, "could not reject unsupported native MCP request"); - } - } + if !self.can_admit_request() { + drop(terminal_tx.send(http::rpc_error( + http_id, + -32000, + "Too many active MCP requests", + ))); + continue; } + let (cancel_tx, cancel_rx) = tokio::sync::oneshot::channel(); + self.active.insert( + request_id.clone(), + ActiveRequest { + server_id: server_id.clone(), + http_id: http_id.clone(), + method: method.clone(), + response_tx: response_tx.clone(), + terminal_tx, + cancel_tx, + }, + ); + let mut tx = self.bridge_tx.clone(); + let cx = connection.clone(); + let request_id_for_task = request_id.clone(); + connection.spawn(async move { + // Dropping the HTTP response stream cancels precisely this ACP request. + let result = tokio::select! { + result = forward_http_request(cx, protocol, server_id, + request_id_for_task, method, params) => Some(result), + () = response_tx.closed() => None, + _ = cancel_rx => None, + }; + tx.send(BridgeMessage::Finished { request_id, result }) + .await + .map_err(agent_client_protocol::Error::into_internal_error)?; + Ok(()) + })?; } - - BridgeMessage::ServerToClientNotification { notification } => { - match self.downstream_mode { - DownstreamMcpMode::Native => { - if let Err(error) = - connection.send_notification_to(Agent, notification.raw) - { - debug!(?error, "could not forward native MCP notification"); - } - } - DownstreamMcpMode::HttpAdapter => { - let connection_id = notification.connection_id; - let Some(active) = self.bridge_connections.get_mut(&connection_id) - else { - debug!( - connection_id, - "ignoring notification for unknown MCP connection" - ); - continue; - }; - let message = UntypedMessage { - method: notification.method, - params: native_params_into_value(notification.params), - }; - if active - .bridge - .try_send(Dispatch::Notification(message)) - .is_some() - { - debug!("discarding MCP notification for unavailable local client"); - } + BridgeMessage::Notification(notification) => { + if self.downstream_mode == DownstreamMcpMode::Native { + connection.send_notification_to(Agent, notification.raw)?; + } else if self.downstream_mode == DownstreamMcpMode::HttpAdapter { + let Some(active) = self.active.get(¬ification.request_id) else { + debug!("dropping notification for stale MCP request"); + continue; + }; + if active.server_id != notification.server_id { + warn!("dropping notification with mismatched MCP server"); + continue; } - DownstreamMcpMode::Unknown | DownstreamMcpMode::Unavailable => { - debug!("ignoring unsupported native MCP notification"); + let mut params = Value::Object(notification.params.unwrap_or_default()); + http::rewrite_subscription_id( + &mut params, + ¬ification.request_id, + &active.http_id, + ); + let message = serde_json::json!({ + "jsonrpc": "2.0", + "method": notification.method, + "params": params, + }); + if active.response_tx.send(message).is_err() { + // Stop only this request. Its final error goes through a + // separate control path that cannot be blocked by a full queue. + let active = self + .active + .remove(¬ification.request_id) + .expect("active request checked above"); + let _ = active.cancel_tx.send(()); + drop(active.terminal_tx.send(http::rpc_error( + active.http_id, + -32000, + "MCP notification queue overflow", + ))); } } } - - BridgeMessage::Disconnected { connection_id } => { - let Some(active) = self.bridge_connections.remove(&connection_id) else { - debug!(connection_id, "local MCP connection was already removed"); - continue; - }; - self.listeners.remove(&active.server_id); - - let Some(protocol) = self.protocol else { - debug!("could not disconnect MCP bridge before ACP initialization"); + BridgeMessage::Finished { request_id, result } => { + let Some(active) = self.active.remove(&request_id) else { continue; }; - let request = protocol.disconnect_request(connection_id)?; - let scheduled = connection - .send_request_to(Client, request) - .on_receiving_result(async move |result| { - match result { - Ok(response) => { - if let Err(error) = - protocol.validate_disconnect_response(response) - { - warn!(?error, "invalid response to mcp/disconnect"); - } - } - Err(error) => { - debug!(?error, "mcp/disconnect failed"); + if let Some(result) = result { + let value = match result { + Ok(mut result) => { + if active.method == "tools/list" { + filter_annotated_tools(&mut result); } + http::rpc_result(active.http_id, &request_id, result) } - Ok(()) - }); - if let Err(error) = scheduled { - debug!( - ?error, - "could not schedule mcp/disconnect response handling" - ); + Err(error) => http::rpc_acp_error(active.http_id, error), + }; + drop(active.terminal_tx.send(value)); } } } @@ -783,286 +522,289 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { } } -fn local_mcp_notification( - protocol: PolyfillProtocol, - connection_id: String, - message: UntypedMessage, -) -> Result, agent_client_protocol::Error> { - if is_cancel_request_notification(&message) { - return Ok(None); +impl BridgeRunner { + fn can_admit_request(&self) -> bool { + self.active.len() < MAX_ACTIVE_REQUESTS + } + + async fn transform_servers( + &mut self, + connection: &ConnectionTo, + servers: Vec, + ) -> Result, agent_client_protocol::Error> { + let protocol = self + .protocol + .ok_or_else(agent_client_protocol::Error::invalid_request)?; + let mut transformed = Vec::with_capacity(servers.len()); + for server in servers { + let Some(native) = protocol.native_server(server.clone()) else { + transformed.push(server); + continue; + }; + match self.downstream_mode { + DownstreamMcpMode::Native => transformed.push(server), + DownstreamMcpMode::HttpAdapter => { + if !self.listeners.contains_key(&native.server_id) { + if self.listeners.len() >= MAX_LISTENERS { + return Err(agent_client_protocol::Error::invalid_params() + .data("too many MCP HTTP listeners")); + } + let listener = TcpListener::bind("127.0.0.1:0") + .await + .map_err(agent_client_protocol::Error::into_internal_error)?; + let port = listener + .local_addr() + .map_err(agent_client_protocol::Error::into_internal_error)? + .port(); + let token = uuid::Uuid::new_v4().simple().to_string() + + &uuid::Uuid::new_v4().simple().to_string(); + connection.spawn(http::run_http_listener( + listener, + native.server_id.clone(), + token.clone(), + self.bridge_tx.clone(), + ))?; + self.listeners.insert( + native.server_id.clone(), + BridgeListener { + tcp_port: port, + token, + }, + ); + } + transformed.push( + self.listeners + .get(&native.server_id) + .expect("listener created") + .declaration(protocol, native)?, + ); + } + DownstreamMcpMode::Unknown | DownstreamMcpMode::Unavailable => { + return Err(agent_client_protocol::Error::invalid_params().data( + "the downstream agent supports neither native nor HTTP MCP transport", + )); + } + } + } + Ok(transformed) } - protocol - .message_notification(connection_id, message) - .map(Some) } -fn reject_native_servers( +/// For each tools/call POST, inspect the current tool schema in that request's +/// scope. This adds an ACP tools/list lookup, but requires no client-side +/// discovery handshake and cannot silently omit an annotated parameter header. +async fn forward_http_request( + connection: ConnectionTo, protocol: PolyfillProtocol, - servers: Vec, - reason: &'static str, -) -> Result, agent_client_protocol::Error> { - if servers - .iter() - .any(|server| protocol.native_server(server.clone()).is_some()) - { - Err(agent_client_protocol::Error::invalid_params().data(reason)) - } else { - Ok(servers) + server_id: String, + request_id: String, + method: String, + params: Option>, +) -> Result { + if method == "tools/call" { + let name = params + .as_ref() + .and_then(|p| p.get("name")) + .and_then(Value::as_str) + .ok_or_else(agent_client_protocol::Error::invalid_params)?; + let meta = params.as_ref().and_then(|p| p.get("_meta")).cloned(); + let mut cursor: Option = None; + let mut seen = HashSet::new(); + loop { + let mut list_params = serde_json::Map::new(); + if let Some(meta) = &meta { + list_params.insert("_meta".into(), meta.clone()); + } + if let Some(cursor) = &cursor { + list_params.insert("cursor".into(), Value::String(cursor.clone())); + } + let lookup = protocol.message_request( + server_id.clone(), + uuid::Uuid::new_v4().to_string(), + "tools/list".into(), + Some(list_params), + None, + )?; + let listing = connection + .send_request_to(Client, lookup) + .block_task() + .await?; + let tools = listing + .get("tools") + .and_then(Value::as_array) + .ok_or_else(|| { + agent_client_protocol::Error::invalid_params() + .data("tools/list result must contain a tools array") + })?; + if let Some(tool) = tools + .iter() + .find(|tool| tool.get("name").and_then(Value::as_str) == Some(name)) + { + if tool + .get("inputSchema") + .is_none_or(|schema| !schema.is_object() || contains_header_annotation(schema)) + { + return Err(agent_client_protocol::Error::invalid_params() + .data("tool uses x-mcp-header or has no verifiable input schema")); + } + break; + } + let Some(next) = listing.get("nextCursor").and_then(Value::as_str) else { + return Err(agent_client_protocol::Error::invalid_params() + .data("tool was not found in tools/list")); + }; + if !seen.insert(next.to_owned()) || seen.len() > 128 { + return Err(agent_client_protocol::Error::invalid_params() + .data("tools/list pagination did not terminate")); + } + cursor = Some(next.to_owned()); + } } + let request = protocol.message_request(server_id, request_id, method, params, None)?; + connection + .send_request_to(Client, request) + .block_task() + .await } -fn respond_unknown_connection(responder: Responder, connection_id: &str) { - let error = agent_client_protocol::Error::invalid_params().data(serde_json::json!({ - "reason": "unknown MCP connection", - "connectionId": connection_id, - })); - if let Err(send_error) = responder.respond_with_error(error) { - debug!( - ?send_error, - connection_id, "could not reject unknown MCP connection" - ); +fn contains_header_annotation(value: &Value) -> bool { + match value { + Value::Object(object) => { + object.contains_key("x-mcp-header") || object.values().any(contains_header_annotation) + } + Value::Array(values) => values.iter().any(contains_header_annotation), + _ => false, } } -#[cfg(test)] -mod tests { - use std::collections::HashMap; - - use agent_client_protocol::{ - Conductor, Dispatch, ErrorCode, Proxy, UntypedMessage, - schema::v1::{ - McpServer, McpServerAcp, McpServerHttp, MessageMcpNotification, MessageMcpRequest, - }, +fn filter_annotated_tools(result: &mut Value) { + let Some(tools) = result.get_mut("tools").and_then(Value::as_array_mut) else { + return; }; - use futures::{StreamExt, channel::mpsc}; + tools.retain(|tool| { + let Some(name) = tool.get("name").and_then(Value::as_str) else { + return false; + }; + if tool + .get("inputSchema") + .is_none_or(|schema| !schema.is_object() || contains_header_annotation(schema)) + { + warn!( + tool = name, + "excluding tool with unsupported x-mcp-header annotation" + ); + return false; + } + true + }); +} - use super::{ - ActiveBridgeConnection, BridgeConnection, BridgeListener, BridgeListeners, BridgeRunner, - DownstreamMcpMode, PolyfillHandler, PolyfillProtocol, local_mcp_notification, - reject_native_servers, - }; +#[cfg(test)] +mod http_limits_tests { + use super::*; #[test] - fn http_declarations_reuse_endpoint_but_preserve_name_and_meta() { - let listener = BridgeListener { tcp_port: 4321 }; - let first_meta = serde_json::Map::from_iter([("source".into(), "first".into())]); - let second_meta = serde_json::Map::from_iter([("source".into(), "second".into())]); - - let first = PolyfillProtocol::V1 - .native_server( - serde_json::to_value(McpServer::Acp( - McpServerAcp::new("first", "shared").meta(first_meta.clone()), - )) - .unwrap(), - ) - .unwrap(); - let second = PolyfillProtocol::V1 - .native_server( - serde_json::to_value(McpServer::Acp( - McpServerAcp::new("second", "shared").meta(second_meta.clone()), - )) - .unwrap(), - ) - .unwrap(); - let first: McpServer = - serde_json::from_value(listener.declaration(PolyfillProtocol::V1, first).unwrap()) - .unwrap(); - let second: McpServer = - serde_json::from_value(listener.declaration(PolyfillProtocol::V1, second).unwrap()) - .unwrap(); - - let McpServer::Http(first) = first else { - panic!("expected HTTP declaration") + fn slow_reader_overflows_by_count_without_blocking_other_requests() { + let (tx, mut rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); + let sender = StreamSender { + tx, + used: Arc::new(AtomicUsize::new(0)), }; - let McpServer::Http(second) = second else { - panic!("expected HTTP declaration") + let (other_tx, mut other_rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); + let other = StreamSender { + tx: other_tx, + used: Arc::new(AtomicUsize::new(0)), }; - assert_eq!(first.url, "http://127.0.0.1:4321"); - assert_eq!(second.url, first.url); - assert_eq!(first.name, "first"); - assert_eq!(second.name, "second"); - assert_eq!(first.meta, Some(first_meta)); - assert_eq!(second.meta, Some(second_meta)); - } - - #[test] - fn downstream_mode_prefers_native_then_http_adaptation() { - assert_eq!( - DownstreamMcpMode::from_capabilities(true, true), - DownstreamMcpMode::Native - ); - assert_eq!( - DownstreamMcpMode::from_capabilities(false, true), - DownstreamMcpMode::Native + for i in 0..MAX_QUEUED_NOTIFICATIONS { + assert!(sender.send(serde_json::json!({"sequence":i})).is_ok()); + } + assert!( + sender + .send(serde_json::json!({"sequence":"overflow"})) + .is_err() ); - assert_eq!( - DownstreamMcpMode::from_capabilities(true, false), - DownstreamMcpMode::HttpAdapter + assert!( + other + .send(serde_json::json!({"sequence":"unaffected"})) + .is_ok() ); - assert_eq!( - DownstreamMcpMode::from_capabilities(false, false), - DownstreamMcpMode::Unavailable + assert_eq!(other_rx.try_recv().unwrap().value["sequence"], "unaffected"); + while rx.try_recv().is_ok() {} + assert_eq!(sender.used.load(Ordering::Relaxed), 0); + assert!( + sender + .send(serde_json::json!({"sequence":"recovered"})) + .is_ok() ); } #[test] - fn local_cancellation_is_not_tunneled_as_an_mcp_message() { - let cancellation = UntypedMessage { - method: "$/cancel_request".to_string(), - params: serde_json::json!({ - "requestId": "loopback-request" - }), + fn large_notification_exceeds_byte_budget_without_reserving_memory() { + let (tx, _rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); + let sender = StreamSender { + tx, + used: Arc::new(AtomicUsize::new(0)), }; - assert_eq!( - local_mcp_notification( - PolyfillProtocol::V1, - "native-connection".to_string(), - cancellation, - ) - .expect("cancellation filtering should not fail"), - None - ); - - let notification = UntypedMessage { - method: "notifications/progress".to_string(), - params: serde_json::json!({ - "progressToken": "token", - "progress": 0.5 - }), - }; - let wrapped = local_mcp_notification( - PolyfillProtocol::V1, - "native-connection".to_string(), - notification, - ) - .expect("the notification should serialize") - .expect("ordinary MCP notifications should be forwarded"); - assert_eq!(wrapped.method, "mcp/message"); - assert_eq!( - wrapped.params["connectionId"], - serde_json::json!("native-connection") - ); - assert_eq!( - wrapped.params["method"], - serde_json::json!("notifications/progress") + assert!( + sender + .send(serde_json::json!({"data":"x".repeat(MAX_QUEUED_BYTES)})) + .is_err() ); + assert_eq!(sender.used.load(Ordering::Relaxed), 0); + assert!(sender.send(serde_json::json!({"data":"ok"})).is_ok()); } #[test] - fn unavailable_mode_rejects_only_native_declarations() { - let standard = vec![ - serde_json::to_value(McpServer::Http(McpServerHttp::new( - "remote", - "https://example.com/mcp", - ))) - .unwrap(), - ]; - assert_eq!( - reject_native_servers(PolyfillProtocol::V1, standard.clone(), "unsupported").unwrap(), - standard - ); - - let error = reject_native_servers( - PolyfillProtocol::V1, - vec![ - serde_json::to_value(McpServer::Acp(McpServerAcp::new("native", "server-1"))) - .unwrap(), - ], - "unsupported", - ) - .expect_err("native declarations require a downstream transport"); - assert_eq!(error.code, ErrorCode::InvalidParams); - assert_eq!(error.data, Some(serde_json::json!("unsupported"))); + fn admission_reopens_when_an_active_request_finishes() { + let (bridge_tx, bridge_rx) = mpsc::channel(1); + let mut runner = BridgeRunner { + bridge_tx, + bridge_rx, + protocol: None, + downstream_mode: DownstreamMcpMode::Unknown, + listeners: HashMap::new(), + active: HashMap::new(), + }; + let (tx, _rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); + let sender = StreamSender { + tx, + used: Arc::new(AtomicUsize::new(0)), + }; + for index in 0..MAX_ACTIVE_REQUESTS { + let (terminal_tx, _terminal_rx) = tokio::sync::oneshot::channel(); + let (cancel_tx, _cancel_rx) = tokio::sync::oneshot::channel(); + runner.active.insert( + index.to_string(), + ActiveRequest { + server_id: String::new(), + http_id: Value::Null, + method: String::new(), + response_tx: sender.clone(), + terminal_tx, + cancel_tx, + }, + ); + } + assert!(!runner.can_admit_request()); + runner.active.remove("0"); + assert!(runner.can_admit_request()); } +} - #[tokio::test(flavor = "current_thread")] - async fn reverse_messages_route_without_stopping_on_unknown_connections() - -> Result<(), agent_client_protocol::Error> { - let known_connection_id = "known-connection"; - let (bridge_tx, bridge_rx) = mpsc::channel(16); - let (to_mcp_client_tx, mut to_mcp_client_rx) = mpsc::channel(16); - let bridge_connections = HashMap::from([( - known_connection_id.to_string(), - ActiveBridgeConnection { - server_id: "test-server".to_string(), - bridge: BridgeConnection::new(to_mcp_client_tx), - }, - )]); - - let proxy = Proxy - .builder() - .with_runner(BridgeRunner { - bridge_tx: bridge_tx.clone(), - bridge_rx, - protocol: Some(PolyfillProtocol::V1), - downstream_mode: DownstreamMcpMode::HttpAdapter, - listeners: BridgeListeners::default(), - bridge_connections, - }) - .with_handler(PolyfillHandler { - protocol: Some(PolyfillProtocol::V1), - bridge_tx, - }); - - Conductor - .builder() - .connect_with(proxy, async move |connection| { - let request_params = serde_json::Map::from_iter([( - "cursor".to_string(), - serde_json::json!("next-page"), - )]); - let request = MessageMcpRequest::new(known_connection_id, "tools/list") - .params(request_params.clone()); - let pending_response = connection.send_request(request); - - let Some(Dispatch::Request(message, responder)) = to_mcp_client_rx.next().await - else { - panic!("expected the request to reach the stored bridge connection") - }; - assert_eq!(message.method, "tools/list"); - assert_eq!(message.params, serde_json::Value::Object(request_params)); - - let inner_response = serde_json::json!({"tools": [{"name": "echo"}]}); - responder.respond(inner_response.clone())?; - let response = pending_response.block_task().await?; - let response: serde_json::Value = serde_json::from_str(response.0.get())?; - assert_eq!(response, inner_response); - - let unknown_error = connection - .send_request(MessageMcpRequest::new( - "missing-connection", - "resources/list", - )) - .block_task() - .await - .expect_err("an unknown connection must receive an error response"); - assert_eq!(unknown_error.code, ErrorCode::InvalidParams); - assert_eq!( - unknown_error.data, - Some(serde_json::json!({ - "reason": "unknown MCP connection", - "connectionId": "missing-connection", - })) - ); - - connection.send_notification(MessageMcpNotification::new( - "missing-connection", - "notifications/progress", - ))?; - connection.send_notification(MessageMcpNotification::new( - known_connection_id, - "notifications/tools/list_changed", - ))?; - - let Some(Dispatch::Notification(notification)) = to_mcp_client_rx.next().await - else { - panic!("expected the known notification after ignoring the unknown one") - }; - assert_eq!(notification.method, "notifications/tools/list_changed"); - assert_eq!(notification.params, serde_json::Value::Null); +#[cfg(test)] +mod tests { + use super::*; - Ok(()) - }) - .await + #[test] + fn annotated_tools_are_not_advertised_or_callable() { + let mut result = serde_json::json!({"tools":[ + {"name":"plain","inputSchema":{"type":"object","properties":{}}}, + {"name":"annotated","inputSchema":{"properties":{"nested":{"properties":{ + "region":{"type":"string","x-mcp-header":"Region"} + }}}}} + ]}); + filter_annotated_tools(&mut result); + assert_eq!(result["tools"].as_array().unwrap().len(), 1); + assert_eq!(result["tools"][0]["name"], "plain"); } } diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs index b3cedc1b..d6d7aea5 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs @@ -2,22 +2,16 @@ use agent_client_protocol::{ Error, JsonRpcMessage, JsonRpcResponse, UntypedMessage, schema::{ InitializeProxyRequest, METHOD_INITIALIZE_PROXY, ProtocolVersion, - v1::{ - ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, DisconnectMcpResponse, - LoadSessionRequest, McpServer, MessageMcpNotification, MessageMcpRequest, - NewSessionRequest, ResumeSessionRequest, - }, + v1::{LoadSessionRequest, McpServer, NewSessionRequest, ResumeSessionRequest}, }, }; use serde_json::{Map, Value}; -#[cfg(feature = "unstable_protocol_v2")] -use agent_client_protocol::schema::v2; - #[cfg(feature = "unstable_session_fork")] use agent_client_protocol::schema::v1::ForkSessionRequest; +#[cfg(feature = "unstable_protocol_v2")] +use agent_client_protocol::schema::v2; -/// ACP schema selected by the conductor's proxy initialization request. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(crate) enum PolyfillProtocol { V1, @@ -28,20 +22,18 @@ pub(crate) enum PolyfillProtocol { impl PolyfillProtocol { pub(crate) fn from_initialize_request(request: &UntypedMessage) -> Result { if request.method() != METHOD_INITIALIZE_PROXY { - return Err(Error::invalid_request() - .data(format!("expected `{METHOD_INITIALIZE_PROXY}` request"))); + return Err(Error::invalid_request().data("expected initialize proxy request")); } - - let requested = request - .params() - .get("protocolVersion") - .cloned() - .ok_or_else(invalid_initialize_protocol_version) - .and_then(|version| { - serde_json::from_value::(version) - .map_err(|_| invalid_initialize_protocol_version()) - })?; - + let requested = serde_json::from_value::( + request + .params() + .get("protocolVersion") + .cloned() + .ok_or_else(|| { + Error::invalid_params().data("missing initialize.protocolVersion") + })?, + ) + .map_err(Error::into_internal_error)?; let protocol = if requested == ProtocolVersion::V1 { Self::V1 } else { @@ -50,22 +42,17 @@ impl PolyfillProtocol { if requested == ProtocolVersion::V2 { Self::V2 } else { - return Err(unsupported_protocol_version(requested)); + return Err(Error::invalid_request() + .data(format!("unsupported ACP protocol version {requested}"))); } } - #[cfg(not(feature = "unstable_protocol_v2"))] { - return Err(unsupported_protocol_version(requested)); + return Err(Error::invalid_request() + .data(format!("unsupported ACP protocol version {requested}"))); } }; - - protocol.validate_initialize_request(request)?; - Ok(protocol) - } - - fn validate_initialize_request(self, request: &UntypedMessage) -> Result<(), Error> { - match self { + match protocol { Self::V1 => { InitializeProxyRequest::parse_message(request.method(), request.params())?; } @@ -74,7 +61,7 @@ impl PolyfillProtocol { v2::InitializeProxyRequest::parse_message(request.method(), request.params())?; } } - Ok(()) + Ok(protocol) } pub(crate) fn transform_initialize_response( @@ -83,72 +70,56 @@ impl PolyfillProtocol { ) -> Result { let mode = match self { Self::V1 => { - let response = agent_client_protocol::schema::v1::InitializeResponse::from_value( + let parsed = agent_client_protocol::schema::v1::InitializeResponse::from_value( "initialize", response.clone(), )?; DownstreamMcpMode::from_capabilities( - response.agent_capabilities.mcp_capabilities.http, - response.agent_capabilities.mcp_capabilities.acp, + parsed.agent_capabilities.mcp_capabilities.http, + parsed.agent_capabilities.mcp_capabilities.acp, ) } #[cfg(feature = "unstable_protocol_v2")] Self::V2 => { - let response = v2::InitializeResponse::from_value("initialize", response.clone())?; - let mcp = response + let parsed = v2::InitializeResponse::from_value("initialize", response.clone())?; + let mcp = parsed .capabilities .session .as_ref() - .and_then(|session| session.mcp.as_ref()); + .and_then(|s| s.mcp.as_ref()); DownstreamMcpMode::from_capabilities( - mcp.is_some_and(|mcp| mcp.http.is_some()), - mcp.is_some_and(|mcp| mcp.acp.is_some()), + mcp.is_some_and(|m| m.http.is_some()), + mcp.is_some_and(|m| m.acp.is_some()), ) } }; - if mode == DownstreamMcpMode::HttpAdapter { - self.advertise_native_mcp(response)?; - } - Ok(mode) - } - - fn advertise_native_mcp(self, response: &mut Value) -> Result<(), Error> { - let response = response - .as_object_mut() - .ok_or_else(|| invalid_initialize_response("result must be an object"))?; - match self { - Self::V1 => { - let mcp = response - .get_mut("agentCapabilities") - .and_then(Value::as_object_mut) - .and_then(|capabilities| capabilities.get_mut("mcpCapabilities")) - .and_then(Value::as_object_mut) - .ok_or_else(|| { - invalid_initialize_response( - "HTTP MCP support did not have an object capability container", - ) - })?; - mcp.insert("acp".into(), Value::Bool(true)); - } - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => { - let mcp = response - .get_mut("capabilities") - .and_then(Value::as_object_mut) - .and_then(|capabilities| capabilities.get_mut("session")) - .and_then(Value::as_object_mut) - .and_then(|session| session.get_mut("mcp")) - .and_then(Value::as_object_mut) - .ok_or_else(|| { - invalid_initialize_response( - "HTTP MCP support did not have an object capability container", - ) - })?; - mcp.insert("acp".into(), Value::Object(Map::new())); + let root = response.as_object_mut().ok_or_else(Error::invalid_params)?; + match self { + Self::V1 => { + let mcp = root + .get_mut("agentCapabilities") + .and_then(Value::as_object_mut) + .and_then(|capabilities| capabilities.get_mut("mcpCapabilities")) + .and_then(Value::as_object_mut) + .ok_or_else(Error::invalid_params)?; + mcp.insert("acp".into(), Value::Bool(true)); + } + #[cfg(feature = "unstable_protocol_v2")] + Self::V2 => { + let mcp = root + .get_mut("capabilities") + .and_then(Value::as_object_mut) + .and_then(|capabilities| capabilities.get_mut("session")) + .and_then(Value::as_object_mut) + .and_then(|session| session.get_mut("mcp")) + .and_then(Value::as_object_mut) + .ok_or_else(Error::invalid_params)?; + mcp.insert("acp".into(), Value::Object(Map::new())); + } } } - Ok(()) + Ok(mode) } pub(crate) fn is_session_setup_method(self, method: &str) -> bool { @@ -184,7 +155,7 @@ impl PolyfillProtocol { "session/fork" => { ForkSessionRequest::parse_message(request.method(), request.params())?; } - method => return Err(unexpected_session_setup_method(method)), + _ => return Err(Error::invalid_request().data("not a session setup method")), }, #[cfg(feature = "unstable_protocol_v2")] Self::V2 => match request.method() { @@ -198,7 +169,7 @@ impl PolyfillProtocol { "session/fork" => { v2::ForkSessionRequest::parse_message(request.method(), request.params())?; } - method => return Err(unexpected_session_setup_method(method)), + _ => return Err(Error::invalid_request().data("not a session setup method")), }, } Ok(()) @@ -206,165 +177,102 @@ impl PolyfillProtocol { pub(crate) fn native_server(self, value: Value) -> Option { let raw = value.as_object()?.clone(); - match self { + let (name, server_id) = match self { Self::V1 => { let McpServer::Acp(server) = serde_json::from_value(value).ok()? else { return None; }; - Some(NativeServer { - raw, - name: server.name, - server_id: server.server_id.to_string(), - }) + (server.name, server.server_id.to_string()) } #[cfg(feature = "unstable_protocol_v2")] Self::V2 => { let v2::McpServer::Acp(server) = serde_json::from_value(value).ok()? else { return None; }; - Some(NativeServer { - raw, - name: server.name, - server_id: server.server_id.to_string(), - }) + (server.name, server.server_id.to_string()) } - } - } - - pub(crate) fn connect_request(self, server_id: String) -> Result { - match self { - Self::V1 => ConnectMcpRequest::new(server_id).to_untyped_message(), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::ConnectMcpRequest::new(server_id).to_untyped_message(), - } - } - - pub(crate) fn connect_response_id(self, response: Value) -> Result { - match self { - Self::V1 => ConnectMcpResponse::from_value("mcp/connect", response) - .map(|response| response.connection_id.to_string()), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::ConnectMcpResponse::from_value("mcp/connect", response) - .map(|response| response.connection_id.to_string()), - } + }; + Some(NativeServer { + raw, + name, + server_id, + }) } pub(crate) fn message_request( self, - connection_id: String, - message: UntypedMessage, + server_id: String, + request_id: String, + method: String, + params: Option>, + meta: Option, ) -> Result { - let (method, params) = message.into_parts(); - let params = into_mcp_params(params)?; - match self { - Self::V1 => MessageMcpRequest::new(connection_id, method) - .params(params) - .to_untyped_message(), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::MessageMcpRequest::new(connection_id, method) - .params(params) - .to_untyped_message(), + let mut wrapper = Map::new(); + wrapper.insert("serverId".into(), server_id.into()); + wrapper.insert("requestId".into(), request_id.into()); + wrapper.insert("method".into(), method.into()); + if let Some(params) = params { + wrapper.insert("params".into(), Value::Object(params)); } - } - - pub(crate) fn message_notification( - self, - connection_id: String, - message: UntypedMessage, - ) -> Result { - let (method, params) = message.into_parts(); - let params = into_mcp_params(params)?; - match self { - Self::V1 => MessageMcpNotification::new(connection_id, method) - .params(params) - .to_untyped_message(), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::MessageMcpNotification::new(connection_id, method) - .params(params) - .to_untyped_message(), + if let Some(meta) = meta { + wrapper.insert("_meta".into(), meta); } - } - - pub(crate) fn parse_message_request( - self, - request: UntypedMessage, - ) -> Result { - match self { - Self::V1 => { - let parsed = MessageMcpRequest::parse_message(request.method(), request.params())?; - Ok(NativeMcpMessage { - raw: request, - connection_id: parsed.connection_id.to_string(), - method: parsed.method, - params: parsed.params, - }) - } - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => { - let parsed = - v2::MessageMcpRequest::parse_message(request.method(), request.params())?; - Ok(NativeMcpMessage { - raw: request, - connection_id: parsed.connection_id.to_string(), - method: parsed.method, - params: parsed.params, - }) - } - } - } - - pub(crate) fn parse_message_notification( - self, - notification: UntypedMessage, - ) -> Result { + let request = UntypedMessage { + method: "mcp/message".into(), + params: Value::Object(wrapper), + }; + // Validate the selected schema without losing unknown wrapper fields. match self { Self::V1 => { - let parsed = MessageMcpNotification::parse_message( - notification.method(), - notification.params(), + agent_client_protocol::schema::v1::MessageMcpRequest::parse_message( + request.method(), + request.params(), )?; - Ok(NativeMcpMessage { - raw: notification, - connection_id: parsed.connection_id.to_string(), - method: parsed.method, - params: parsed.params, - }) } #[cfg(feature = "unstable_protocol_v2")] Self::V2 => { - let parsed = v2::MessageMcpNotification::parse_message( - notification.method(), - notification.params(), - )?; - Ok(NativeMcpMessage { - raw: notification, - connection_id: parsed.connection_id.to_string(), - method: parsed.method, - params: parsed.params, - }) + v2::MessageMcpRequest::parse_message(request.method(), request.params())?; } } + Ok(request) } - pub(crate) fn disconnect_request(self, connection_id: String) -> Result { - match self { - Self::V1 => DisconnectMcpRequest::new(connection_id).to_untyped_message(), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::DisconnectMcpRequest::new(connection_id).to_untyped_message(), - } - } - - pub(crate) fn validate_disconnect_response(self, response: Value) -> Result<(), Error> { - match self { + pub(crate) fn parse_notification( + self, + raw: UntypedMessage, + ) -> Result { + let (server_id, request_id, method, params) = match self { Self::V1 => { - DisconnectMcpResponse::from_value("mcp/disconnect", response)?; + let parsed = + agent_client_protocol::schema::v1::MessageMcpNotification::parse_message( + raw.method(), + raw.params(), + )?; + ( + parsed.server_id.to_string(), + parsed.request_id.to_string(), + parsed.method, + parsed.params, + ) } #[cfg(feature = "unstable_protocol_v2")] Self::V2 => { - v2::DisconnectMcpResponse::from_value("mcp/disconnect", response)?; + let parsed = v2::MessageMcpNotification::parse_message(raw.method(), raw.params())?; + ( + parsed.server_id.to_string(), + parsed.request_id.to_string(), + parsed.method, + parsed.params, + ) } - } - Ok(()) + }; + Ok(NativeMcpNotification { + raw, + server_id, + request_id, + method, + params, + }) } } @@ -401,16 +309,19 @@ impl NativeServer { mut self, protocol: PolyfillProtocol, url: String, + token: &str, ) -> Result { self.raw.remove("serverId"); - self.raw.insert("type".into(), Value::String("http".into())); - self.raw.insert("name".into(), Value::String(self.name)); - self.raw.insert("url".into(), Value::String(url)); - // V1 requires the field and v2 accepts it. Keeping the explicit empty - // list gives both versions one stable raw compatibility shape. - self.raw.insert("headers".into(), Value::Array(Vec::new())); + self.raw.insert("type".into(), "http".into()); + self.raw.insert("name".into(), self.name.into()); + self.raw.insert("url".into(), url.into()); + self.raw.insert( + "headers".into(), + serde_json::json!([ + { "name": "Authorization", "value": format!("Bearer {token}") } + ]), + ); let declaration = Value::Object(self.raw); - match protocol { PolyfillProtocol::V1 => { serde_json::from_value::(declaration.clone()) @@ -426,272 +337,10 @@ impl NativeServer { } } -#[derive(Debug)] -pub(crate) struct NativeMcpMessage { +pub(crate) struct NativeMcpNotification { pub(crate) raw: UntypedMessage, - pub(crate) connection_id: String, + pub(crate) server_id: String, + pub(crate) request_id: String, pub(crate) method: String, pub(crate) params: Option>, } - -pub(crate) fn native_params_into_value(params: Option>) -> Value { - params.map_or(Value::Null, Value::Object) -} - -fn into_mcp_params(params: Value) -> Result>, Error> { - match params { - Value::Null => Ok(None), - Value::Object(params) => Ok(Some(params)), - params => Err(Error::invalid_params().data(serde_json::json!({ - "reason": "MCP message params must be an object or null", - "params": params, - }))), - } -} - -fn invalid_initialize_protocol_version() -> Error { - Error::invalid_params().data("initialize.protocolVersion must be a valid ACP protocol version") -} - -fn unsupported_protocol_version(version: ProtocolVersion) -> Error { - Error::invalid_request().data(format!( - "MCP-over-ACP polyfill does not support ACP protocol version {version}" - )) -} - -fn unexpected_session_setup_method(method: &str) -> Error { - Error::invalid_request().data(format!( - "`{method}` is not a session setup method for the selected ACP version" - )) -} - -fn invalid_initialize_response(reason: &'static str) -> Error { - Error::invalid_params().data(format!("invalid initialize response: {reason}")) -} - -#[cfg(test)] -mod tests { - use agent_client_protocol::{ - JsonRpcMessage, - schema::{ProtocolVersion, v1}, - }; - - #[cfg(feature = "unstable_protocol_v2")] - use agent_client_protocol::{ErrorCode, JsonRpcResponse}; - - use super::PolyfillProtocol; - - #[test] - fn http_declaration_preserves_extension_fields() { - let declaration = serde_json::json!({ - "type": "acp", - "name": "native", - "serverId": "native-id", - "_meta": { - "source": "test" - }, - "futureField": { - "preserve": true - } - }); - let native = PolyfillProtocol::V1 - .native_server(declaration) - .expect("the declaration should be recognized as native MCP"); - - let transformed = native - .http_declaration(PolyfillProtocol::V1, "http://127.0.0.1:4321".to_string()) - .expect("the transformed declaration should be valid v1 MCP"); - - assert_eq!( - transformed, - serde_json::json!({ - "type": "http", - "name": "native", - "url": "http://127.0.0.1:4321", - "headers": [], - "_meta": { - "source": "test" - }, - "futureField": { - "preserve": true - } - }) - ); - } - - #[test] - fn native_message_keeps_the_original_wrapper() { - let request = agent_client_protocol::UntypedMessage { - method: "mcp/message".to_string(), - params: serde_json::json!({ - "connectionId": "connection", - "method": "tools/list", - "params": { - "cursor": "next" - }, - "_meta": { - "trace": "preserve" - }, - "futureField": true - }), - }; - - let parsed = PolyfillProtocol::V1 - .parse_message_request(request.clone()) - .expect("the native wrapper should parse"); - - assert_eq!(parsed.raw, request); - assert_eq!(parsed.connection_id, "connection"); - assert_eq!(parsed.method, "tools/list"); - assert_eq!( - parsed.params, - Some(serde_json::Map::from_iter([( - "cursor".to_string(), - serde_json::json!("next") - )])) - ); - } - - #[test] - fn v1_session_setup_methods_match_the_stable_schema() { - assert!(PolyfillProtocol::V1.is_session_setup_method("session/new")); - assert!(PolyfillProtocol::V1.is_session_setup_method("session/load")); - assert!(PolyfillProtocol::V1.is_session_setup_method("session/resume")); - assert_eq!( - PolyfillProtocol::V1.is_session_setup_method("session/fork"), - cfg!(feature = "unstable_session_fork") - ); - assert!(!PolyfillProtocol::V1.is_session_setup_method("session/prompt")); - } - - #[test] - fn session_setup_validation_allows_extensions_but_rejects_invalid_fields() { - let mut request = v1::NewSessionRequest::new(std::path::PathBuf::from("/tmp")) - .to_untyped_message() - .expect("the session request should serialize"); - request.params["futureField"] = serde_json::json!({ - "preserve": true - }); - PolyfillProtocol::V1 - .validate_session_setup_request(&request) - .expect("extension fields should remain forward-compatible"); - - request.params["cwd"] = serde_json::json!(42); - let error = PolyfillProtocol::V1 - .validate_session_setup_request(&request) - .expect_err("invalid selected-schema fields must be rejected"); - assert_eq!(error.code, agent_client_protocol::ErrorCode::InvalidParams); - } - - #[cfg(feature = "unstable_protocol_v2")] - #[test] - fn v2_session_setup_methods_exclude_v1_load() { - assert!(PolyfillProtocol::V2.is_session_setup_method("session/new")); - assert!(!PolyfillProtocol::V2.is_session_setup_method("session/load")); - assert!(PolyfillProtocol::V2.is_session_setup_method("session/resume")); - assert_eq!( - PolyfillProtocol::V2.is_session_setup_method("session/fork"), - cfg!(feature = "unstable_session_fork") - ); - assert!(!PolyfillProtocol::V2.is_session_setup_method("session/prompt")); - } - - #[cfg(feature = "unstable_protocol_v2")] - #[test] - fn future_protocol_version_is_not_assumed_to_be_v2() { - let initialize = agent_client_protocol::schema::v2::InitializeRequest::new( - ProtocolVersion::V2, - agent_client_protocol::schema::v2::Implementation::new("test", "1.0.0"), - ); - let mut request = - agent_client_protocol::schema::v2::InitializeProxyRequest::new(initialize) - .to_untyped_message() - .expect("the initialize request should serialize"); - request.params["protocolVersion"] = serde_json::json!(3); - - let error = PolyfillProtocol::from_initialize_request(&request) - .expect_err("an unselected future schema must not be interpreted as v2"); - - assert_eq!(error.code, ErrorCode::InvalidRequest); - assert_eq!( - error.data, - Some(serde_json::json!( - "MCP-over-ACP polyfill does not support ACP protocol version 3" - )) - ); - } - - #[cfg(feature = "unstable_protocol_v2")] - #[test] - fn v2_initialize_adaptation_preserves_the_raw_response() { - use agent_client_protocol::schema::v2; - - let response = v2::InitializeResponse::new( - ProtocolVersion::V2, - v2::Implementation::new("test", "1.0.0"), - ) - .capabilities( - v2::AgentCapabilities::new().session( - v2::SessionCapabilities::new() - .mcp(v2::McpCapabilities::new().http(v2::McpHttpCapabilities::new())), - ), - ); - let mut response = serde_json::to_value(response).expect("the response should serialize"); - response["futureField"] = serde_json::json!({ - "preserve": true - }); - - let mode = PolyfillProtocol::V2 - .transform_initialize_response(&mut response) - .expect("the v2 HTTP capability should be adaptable"); - - assert_eq!(mode, super::DownstreamMcpMode::HttpAdapter); - assert_eq!( - response["capabilities"]["session"]["mcp"]["acp"], - serde_json::json!({}) - ); - assert_eq!( - response["futureField"], - serde_json::json!({ - "preserve": true - }) - ); - - v2::InitializeResponse::from_value("initialize", response) - .expect("the adapted response should remain valid v2"); - } - - #[cfg(feature = "unstable_protocol_v2")] - #[test] - fn v2_disconnect_uses_and_validates_the_selected_schema() { - use agent_client_protocol::schema::v2; - - let request = PolyfillProtocol::V2 - .disconnect_request("connection".to_string()) - .expect("the v2 disconnect request should serialize"); - let parsed = v2::DisconnectMcpRequest::parse_message(request.method(), request.params()) - .expect("the disconnect request should be valid v2"); - assert_eq!(parsed.connection_id.to_string(), "connection"); - - let response = serde_json::to_value(v2::DisconnectMcpResponse::new()) - .expect("the v2 disconnect response should serialize"); - PolyfillProtocol::V2 - .validate_disconnect_response(response) - .expect("the v2 disconnect response should validate"); - } - - #[test] - fn v1_initialize_request_selects_v1() { - let request = agent_client_protocol::schema::InitializeProxyRequest { - initialize: v1::InitializeRequest::new(ProtocolVersion::V1), - } - .to_untyped_message() - .expect("the initialize request should serialize"); - - assert_eq!( - PolyfillProtocol::from_initialize_request(&request) - .expect("the request should select v1"), - PolyfillProtocol::V1 - ); - } -} diff --git a/src/agent-client-protocol-rmcp/Cargo.toml b/src/agent-client-protocol-rmcp/Cargo.toml index 4952b98d..ab207164 100644 --- a/src/agent-client-protocol-rmcp/Cargo.toml +++ b/src/agent-client-protocol-rmcp/Cargo.toml @@ -13,11 +13,16 @@ categories = ["development-tools"] [features] default = [] unstable_mcp_over_acp = ["agent-client-protocol/unstable_mcp_over_acp"] +unstable_protocol_v2 = ["agent-client-protocol/unstable_protocol_v2"] [[example]] name = "with_mcp_server" required-features = ["unstable_mcp_over_acp"] +[[example]] +name = "stateless_native_mcp" +required-features = ["unstable_mcp_over_acp", "unstable_protocol_v2"] + [dependencies] agent-client-protocol = { workspace = true, features = ["schemars"] } futures.workspace = true diff --git a/src/agent-client-protocol-rmcp/README.md b/src/agent-client-protocol-rmcp/README.md index 4b12c962..7cee740c 100644 --- a/src/agent-client-protocol-rmcp/README.md +++ b/src/agent-client-protocol-rmcp/README.md @@ -9,12 +9,20 @@ runtime-agnostic MCP server framework from `agent-client-protocol`. It lets you Rust, serve them directly, or attach them to an ACP proxy. Attached servers are advertised with the opt-in native MCP-over-ACP transport: -`McpServer::Acp` plus `mcp/connect`, `mcp/message`, and `mcp/disconnect`. This +`McpServer::Acp` plus request-scoped `mcp/message` operations targeting MCP +2026-07-28. There is no MCP initialization or connect/disconnect exchange. This crate does not enable the core SDK's `unstable_mcp_over_acp` feature merely to build or directly serve a server. Enable this crate's matching `unstable_mcp_over_acp` feature when using `with_mcp_server`. Use -`agent-client-protocol-polyfill` when the final agent accepts HTTP but not -ACP-transport MCP servers. +`agent-client-protocol-polyfill` when the final agent has a modern MCP HTTP +client but does not consume ACP-transport MCP servers natively. + +For a direct ACP client/agent example using real rmcp tools, run: + +```sh +cargo run -p agent-client-protocol-rmcp --example stateless_native_mcp \ + --features unstable_mcp_over_acp,unstable_protocol_v2 +``` ## Usage diff --git a/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs b/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs new file mode 100644 index 00000000..81a206fe --- /dev/null +++ b/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs @@ -0,0 +1,123 @@ +//! Run with `cargo run -p agent-client-protocol-rmcp --example stateless_native_mcp +//! --features unstable_mcp_over_acp,unstable_protocol_v2`. +//! No MCP initialize or separate MCP transport: the client attaches an rmcp service +//! to an ACP session and the agent invokes it through `mcp/message`. + +use agent_client_protocol::{ + Agent, Client, Error, Responder, V2ConnectionTo, + mcp_server::McpServer, + schema::{ProtocolVersion, v2}, +}; +use agent_client_protocol_rmcp::McpServerExt; +use rmcp::{ + ErrorData, RoleServer, ServerHandler, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, ServerCapabilities, ServerConfig, + }, + service::RequestContext, +}; +use serde_json::json; +use std::sync::{Arc, Mutex}; +use tokio::sync::oneshot; + +struct Echo; + +impl ServerHandler for Echo { + fn get_info(&self) -> ServerConfig { + ServerConfig::new(ServerCapabilities::builder().enable_tools().build()) + } + + fn call_tool( + &self, + params: CallToolRequestParams, + _cx: RequestContext, + ) -> impl std::future::Future> + Send { + std::future::ready(if params.name == "echo" { + Ok(CallToolResult::structured(json!({"echoed": params.arguments})).into()) + } else { + Err(ErrorData::invalid_params("unknown tool", None)) + }) + } +} + +#[tokio::main] +async fn main() -> Result<(), Error> { + let (done_tx, done_rx) = oneshot::channel(); + let done_tx = Arc::new(Mutex::new(Some(done_tx))); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("echo-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new() + .session(v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + )), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let server_id = match request.mcp_servers.as_slice() { + [v2::McpServer::Acp(server)] if server.name == "echo" => { + server.server_id.clone() + } + other => panic!("unexpected MCP declaration: {other:?}"), + }; + let done_tx = done_tx.lock().unwrap().take().expect("one session"); + let call_cx = cx.clone(); + cx.spawn(async move { + let mut params = json!({"name": "echo", "arguments": {"message": "hello ACP"}}); + params["_meta"] = json!({ + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + "io.modelcontextprotocol/clientInfo": {"name": "echo-agent", "version": "1"} + }); + let response = call_cx + .send_request( + v2::MessageMcpRequest::new(server_id, "echo-1", "tools/call") + .params(params.as_object().expect("object params").clone()), + ) + .block_task() + .await; + drop(done_tx.send(response)); + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new( + "echo-session", + ))) + }, + agent_client_protocol::on_receive_request!(), + ); + + Client + .v2() + .connect_with(agent, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("echo-client", "1"), + )) + .block_task() + .await?; + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(McpServer::::from_rmcp("echo", || Echo))? + .start_session() + .block_task() + .await?; + let response = done_rx.await.map_err(Error::into_internal_error)??; + println!("{}", response.0.get()); + Ok(()) + }) + .await +} diff --git a/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs b/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs new file mode 100644 index 00000000..6ce95acd --- /dev/null +++ b/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs @@ -0,0 +1,320 @@ +//! Native ACP attachment of an rmcp service (not standalone MCP transport). +#![cfg(all(feature = "unstable_protocol_v2", feature = "unstable_mcp_over_acp"))] + +use std::{ + future::Future, + sync::{Arc, Mutex}, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, Client, Error, Responder, V2ConnectionTo, + mcp_server::McpServer, + schema::{ProtocolVersion, v2}, +}; +use agent_client_protocol_rmcp::McpServerExt; +use rmcp::{ + ErrorData, RoleServer, ServerHandler, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, InputRequiredResult, + ServerCapabilities, ServerConfig, SubscriptionFilter, + }, + service::{RequestContext, SubscriptionContext}, +}; +use serde_json::{Value, json}; +use tokio::sync::{mpsc, oneshot}; + +fn meta(marker: &str) -> Value { + json!({"io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {"elicitation": {"form": {}}}, + "io.modelcontextprotocol/clientInfo": {"name": "native-acp", "version": "1"}, + "example/marker": marker}) +} + +async fn message( + cx: &V2ConnectionTo, + server: &v2::McpServerAcpId, + id: &str, + method: &str, + mut params: Value, + marker: &str, +) -> Result { + params["_meta"] = meta(marker); + let response = cx + .send_request( + v2::MessageMcpRequest::new(server.clone(), id.to_owned(), method) + .params(params.as_object().expect("object params").clone()), + ) + .block_task() + .await?; + serde_json::from_str(response.0.get()).map_err(Error::into_internal_error) +} + +struct DropSignal(Arc>>>); +impl Drop for DropSignal { + fn drop(&mut self) { + if let Some(tx) = self.0.lock().unwrap().take() { + let _ = tx.send(()); + } + } +} + +struct Service { + _drop: DropSignal, + started: Arc>>>, + stopped: Arc>>>, +} +impl ServerHandler for Service { + fn get_info(&self) -> ServerConfig { + ServerConfig::new( + ServerCapabilities::builder() + .enable_tools() + .enable_tool_list_changed() + .build(), + ) + } + fn call_tool( + &self, + request: CallToolRequestParams, + cx: RequestContext, + ) -> impl Future> + Send { + std::future::ready(match request.name.as_ref() { + "retry" if request.request_state.is_none() => { + let inputs = serde_json::from_value(json!({"confirmation": { + "method": "elicitation/create", "params": {"mode": "form", + "message": "Confirm", "requestedSchema": {"type": "object", + "properties": {"approved": {"type": "boolean"}}}} + }})) + .expect("valid elicitation"); + Ok(InputRequiredResult::new(Some(inputs), Some("retry-state".into())).into()) + } + "retry" if request.request_state.as_deref() == Some("retry-state") => Ok( + CallToolResult::structured(json!({"marker": cx.meta.get("example/marker"), + "responses": request.input_responses})) + .into(), + ), + "echo" => Ok(CallToolResult::structured( + json!({"marker": cx.meta.get("example/marker")}), + ) + .into()), + _ => Err(ErrorData::invalid_params( + "unknown tool or state", + Some(json!({"source": "rmcp"})), + )), + }) + } + fn accepted_subscription_filter( + &self, + requested: &SubscriptionFilter, + ) -> Option { + Some(requested.clone()) + } + async fn listen(&self, cx: SubscriptionContext) -> Result<(), ErrorData> { + let _stopped = DropSignal(self.stopped.clone()); + cx.sink() + .notify_tool_list_changed() + .await + .map_err(|e| ErrorData::internal_error(e.to_string(), None))?; + if let Some(tx) = self.started.lock().unwrap().take() { + let _ = tx.send(()); + } + cx.cancelled().await; + Ok(()) + } +} + +async fn exercise( + cx: V2ConnectionTo, + server: v2::McpServerAcpId, + started: oneshot::Receiver<()>, + stopped: oneshot::Receiver<()>, + dropped: oneshot::Receiver<()>, +) -> Result { + let direct = message( + &cx, + &server, + "direct-1", + "tools/call", + json!({"name": "echo", "arguments": {}}), + "direct", + ) + .await?; + assert_eq!(direct["structuredContent"]["marker"], "direct"); + let discovered = message( + &cx, + &server, + "discover-1", + "server/discover", + json!({}), + "discover", + ) + .await?; + assert!( + discovered["supportedVersions"] + .as_array() + .unwrap() + .contains(&json!("2026-07-28")) + ); + let first = message( + &cx, + &server, + "retry-1", + "tools/call", + json!({"name": "retry", "arguments": {}}), + "first", + ) + .await?; + assert_eq!(first["resultType"], "input_required"); + assert_eq!( + first["inputRequests"]["confirmation"]["method"], + "elicitation/create" + ); + let responses = json!({"confirmation": {"action": "accept", "content": {"approved": true}}}); + let retry = message( + &cx, + &server, + "retry-2", + "tools/call", + json!({"name": "retry", "arguments": {}, "requestState": first["requestState"], + "inputResponses": responses}), + "second", + ) + .await?; + assert_eq!(retry["structuredContent"]["marker"], "second"); + assert_eq!(retry["structuredContent"]["responses"], responses); + let error = message( + &cx, + &server, + "error-1", + "tools/call", + json!({"name": "missing", "arguments": {}}), + "error", + ) + .await + .expect_err("rmcp error"); + assert_eq!( + serde_json::to_value(error)?["data"], + json!({"source": "rmcp"}) + ); + let mut params = json!({"notifications": {"toolsListChanged": true}}); + params["_meta"] = meta("listen"); + let subscription = cx.send_request( + v2::MessageMcpRequest::new(server.clone(), "listen-1", "subscriptions/listen") + .params(params.as_object().expect("object params").clone()), + ); + started.await.map_err(Error::into_internal_error)?; + let parallel = message( + &cx, + &server, + "parallel-1", + "tools/call", + json!({"name": "echo", "arguments": {}}), + "parallel", + ) + .await?; + assert_eq!(parallel["structuredContent"]["marker"], "parallel"); + subscription.cancel()?; + stopped.await.map_err(Error::into_internal_error)?; + dropped.await.map_err(Error::into_internal_error)?; + Ok(server) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn native_acp_stateless_rmcp_lifecycle() -> Result<(), Error> { + tokio::time::timeout(Duration::from_secs(10), async { + let (start_tx, start_rx) = oneshot::channel(); + let (stop_tx, stop_rx) = oneshot::channel(); + let (drop_tx, drop_rx) = oneshot::channel(); + let (result_tx, result_rx) = oneshot::channel(); + let invocation = Arc::new(Mutex::new(Some((start_rx, stop_rx, drop_rx, result_tx)))); + let (notifications_tx, mut notifications_rx) = mpsc::unbounded_channel(); + let started = Arc::new(Mutex::new(Some(start_tx))); + let stopped = Arc::new(Mutex::new(Some(stop_tx))); + let dropped = Arc::new(Mutex::new(Some(drop_tx))); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("native-rmcp-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + ), + ), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let server = match request.mcp_servers.as_slice() { + [v2::McpServer::Acp(server)] if server.name == "real-rmcp" => { + server.server_id.clone() + } + other => panic!("unexpected declarations: {other:?}"), + }; + let (start_rx, stop_rx, drop_rx, result_tx) = + invocation.lock().unwrap().take().expect("one session"); + let call_cx = cx.clone(); + cx.spawn(async move { + let result = exercise(call_cx, server, start_rx, stop_rx, drop_rx).await; + drop(result_tx.send(result)); + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new( + "native-session", + ))) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_notification( + async move |notification: v2::MessageMcpNotification, + _cx: V2ConnectionTo| { + notifications_tx + .send(notification) + .map_err(Error::into_internal_error) + }, + agent_client_protocol::on_receive_notification!(), + ); + + Client.v2().connect_with(agent, async move |cx| { + cx.send_request(v2::InitializeRequest::new(ProtocolVersion::V2, + v2::Implementation::new("native-rmcp-client", "1"))).block_task().await?; + let server = McpServer::::from_rmcp("real-rmcp", move || Service { + _drop: DropSignal(dropped.clone()), + started: started.clone(), stopped: stopped.clone(), + }); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)?.start_session().block_task().await?; + let server_id = result_rx.await.map_err(Error::into_internal_error)??; + let acknowledgment = notifications_rx.recv().await.expect("acknowledgment"); + let update = notifications_rx.recv().await.expect("filtered update"); + assert_eq!(acknowledgment.method, "notifications/subscriptions/acknowledged"); + assert_eq!( + acknowledgment.params.as_ref().unwrap()["notifications"]["toolsListChanged"], + json!(true), + "the rmcp subscription must accept the requested notification filter" + ); + assert_eq!(update.method, "notifications/tools/list_changed"); + for notification in [acknowledgment, update] { + assert_eq!(notification.server_id, server_id); + assert_eq!(notification.request_id.0.as_ref(), "listen-1"); + assert_eq!(notification.params.as_ref().unwrap()["_meta"] + ["io.modelcontextprotocol/subscriptionId"], json!("listen-1")); + } + Ok(()) + }).await + }) + .await + .expect("native ACP/rmcp operation or cleanup timed out") +} diff --git a/src/agent-client-protocol-test/src/testy.rs b/src/agent-client-protocol-test/src/testy.rs index 9ae4b604..61ca2e9b 100644 --- a/src/agent-client-protocol-test/src/testy.rs +++ b/src/agent-client-protocol-test/src/testy.rs @@ -1464,15 +1464,26 @@ impl Testy { operation: F, ) -> Result where - F: FnOnce(rmcp::service::RunningService) -> Fut, + F: FnOnce( + rmcp::service::RunningService, + ) -> Fut, Fut: std::future::Future>, { use rmcp::{ - ServiceExt, + ClientLifecycleMode, ClientServiceExt, ServiceExt, + model::{ClientCapabilities, ClientConfig, Implementation, ProtocolVersion}, transport::{ConfigureCommandExt, TokioChildProcess}, }; use tokio::process::Command; + let client_config = || { + ClientConfig::new( + ClientCapabilities::default(), + Implementation::new("testy", env!("CARGO_PKG_VERSION")), + ) + .with_protocol_version(ProtocolVersion::V_2026_07_28) + }; + let mcp_servers = self .get_mcp_servers(session_id) .ok_or_else(|| anyhow::anyhow!("Session not found"))?; @@ -1490,7 +1501,9 @@ impl Testy { match mcp_server { McpServer::Stdio(stdio) => { self.run_until_session_cancelled(session_id, async move { - let mcp_client = () + // Standalone stdio servers may still require initialize; + // native-over-ACP HTTP below uses discover without fallback. + let mcp_client = ClientConfig::default() .serve(TokioChildProcess::new( Command::new(&stdio.command).configure(|cmd| { cmd.args(&stdio.args); @@ -1516,9 +1529,14 @@ impl Testy { .custom_headers(http_headers(&http.headers)?); self.run_until_session_cancelled(session_id, async move { - let mcp_client = - ().serve(StreamableHttpClientTransport::from_config(transport_config)) - .await?; + let mcp_client = client_config() + .serve_with_lifecycle( + StreamableHttpClientTransport::from_config(transport_config), + ClientLifecycleMode::Discover { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + }, + ) + .await?; operation(mcp_client).await }) diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index fa0324ee..35f65fac 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -2,6 +2,18 @@ ## [Unreleased] +### Changed (unstable MCP-over-ACP) + +- Target MCP 2026-07-28 with server-addressed `mcp/message` operations and + logical `McpRequestId`s. Remove connect/disconnect and reverse MCP requests; + providers send request-scoped notifications and use ACP cancellation. +- Create an independent backend per operation and expose `request_id()` in + attached MCP contexts instead of `connection_id()`. Preserve standalone MCP + serving independently of the unstable ACP transport feature. +- Validate modern request metadata, restrict discovery to the binding's MCP + revision, and add native admission/payload limits. End-to-end native queue + backpressure remains required before stabilization. + ### Added - Add a default-enabled `schemars` feature that forwards JSON Schema support to diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index bae99ca1..699ae9cf 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -236,7 +236,7 @@ impl Serialize for TransportBatch { } impl TransportFrame { - pub(crate) fn inspect_messages( + fn inspect_messages( &self, observer: &mut impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error>, ) -> Result<(), crate::Error> { diff --git a/src/agent-client-protocol/src/mcp_server/active_session.rs b/src/agent-client-protocol/src/mcp_server/active_session.rs index a20d23d5..a5e1ce0f 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -8,6 +8,7 @@ use futures::{ use serde_json::{Map, Value}; use std::{ collections::HashMap, + io::Write, marker::PhantomData, sync::{Arc, Mutex, Weak}, }; @@ -25,6 +26,13 @@ use crate::{ util::MatchDispatchFrom, }; +// These bound admitted work and individual payloads, not the SDK's underlying +// Channel/outgoing queues. End-to-end native backpressure is separate transport work. +const MAX_ACTIVE_REQUESTS: usize = 64; +const MAX_PAYLOAD_BYTES: usize = 16 * 1024 * 1024; +const MCP_VERSION: &str = "2026-07-28"; +type ActiveRequests = Arc>>>; + pub(super) struct V1McpProtocol; #[cfg(feature = "unstable_protocol_v2")] pub(super) struct V2McpProtocol; @@ -99,7 +107,7 @@ impl McpProtocol for V2McpProtocol { pub(super) struct McpActiveSession { server_id: McpServerAcpId, mcp_connect: Arc>, - active: Arc>>>, + active: ActiveRequests, protocol: PhantomData Protocol>, } @@ -119,6 +127,52 @@ impl Drop for ActiveRequest { } } +fn admit_request( + active: &ActiveRequests, + id: McpRequestId, +) -> Result<(ActiveRequest, oneshot::Receiver<()>), crate::Error> { + let (stop_tx, stop_rx) = oneshot::channel(); + let mut requests = active.lock().expect("MCP request registry poisoned"); + if requests.contains_key(&id) { + return Err(crate::Error::invalid_params().data("duplicate active MCP requestId")); + } + if requests.len() >= MAX_ACTIVE_REQUESTS { + return Err( + crate::Error::new(-32000, "MCP active request limit exceeded") + .data(serde_json::json!({"limit": MAX_ACTIVE_REQUESTS})), + ); + } + requests.insert(id.clone(), stop_tx); + Ok(( + ActiveRequest { + active: Arc::downgrade(active), + id, + }, + stop_rx, + )) +} + +/// Count serialized bytes without allocating another copy of a potentially large payload. +fn check_payload_size(value: &impl serde::Serialize, limit: usize) -> Result<(), crate::Error> { + struct Budget(usize); + impl Write for Budget { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0 = self + .0 + .checked_sub(bytes.len()) + .ok_or_else(|| std::io::Error::other("MCP payload limit exceeded"))?; + Ok(bytes.len()) + } + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + serde_json::to_writer(Budget(limit), value).map_err(|_| { + crate::Error::new(-32000, "MCP payload limit exceeded") + .data(serde_json::json!({"limitBytes": limit})) + }) +} + impl McpActiveSession where Counterpart: HasPeer, @@ -157,31 +211,20 @@ where } let request_id = Protocol::request_id(&request); let (method, params) = Protocol::into_request(request); - if let Err(error) = validate_modern_request(&method, params.as_ref()) { + if let Err(error) = validate_modern_request(&method, params.as_ref()) + .and_then(|()| check_payload_size(&(&method, ¶ms, &request_id), MAX_PAYLOAD_BYTES)) + { responder.respond_with_error(error)?; return Ok(Handled::Yes); } - let (stop_tx, stop_rx) = oneshot::channel(); - let duplicate = { - let mut active = self.active.lock().expect("MCP request registry poisoned"); - if active.contains_key(&request_id) { - true - } else { - active.insert(request_id.clone(), stop_tx); - false + let (guard, stop_rx) = match admit_request(&self.active, request_id.clone()) { + Ok(admitted) => admitted, + Err(error) => { + responder.respond_with_error(error)?; + return Ok(Handled::Yes); } }; - if duplicate { - responder.respond_with_error( - crate::Error::invalid_params().data("duplicate active MCP requestId"), - )?; - return Ok(Handled::Yes); - } - let guard = ActiveRequest { - active: Arc::downgrade(&self.active), - id: request_id.clone(), - }; let backend = self.mcp_connect.connect(McpConnectionTo { context: McpConnectionContext::Acp { server_id: server_id.clone(), @@ -215,6 +258,7 @@ where } let spawn_result = connection.spawn(async move { let inner_id = RequestId::Str(request_id.0.to_string()); + let is_discovery = method == "server/discover"; let process = async { let raw = RawJsonRpcMessage::request( method, @@ -226,30 +270,34 @@ where .unbounded_send(TransportFrame::Single(raw)) .map_err(crate::Error::into_internal_error)?; while let Some(frame) = client.rx.next().await { - let mut result = None; - frame.inspect_messages(&mut |message| { - // A response ends the request, even within a batch. Notifications - // following it must not escape after the operation has completed. - if result.is_some() { - return Ok(()); - } - match message { - RawJsonRpcMessage::Response(response) => { - if message.response_id() != Some(&inner_id) { - return Err(crate::Error::invalid_params() - .data("MCP backend returned a different request ID")); - } - result = Some(match response { - crate::schema::v1::Response::Result { result, .. } => { - Ok(result.clone()) - } - crate::schema::v1::Response::Error { error, .. } => { - Err(error.clone()) + let TransportFrame::Single(message) = frame else { + return Err(crate::Error::invalid_request() + .data("MCP backends must send individual valid JSON-RPC messages")); + }; + if matches!(message, RawJsonRpcMessage::Response(_)) + && message.response_id() != Some(&inner_id) + { + return Err(crate::Error::invalid_params() + .data("MCP backend returned a different request ID")); + } + match message { + RawJsonRpcMessage::Response(response) => { + check_payload_size(&response, MAX_PAYLOAD_BYTES)?; + // Returning ends notification forwarding before the terminal reply. + return match response { + crate::schema::v1::Response::Result { mut result, .. } => { + if is_discovery { + constrain_discovery_versions(&mut result)?; } - }); - } - RawJsonRpcMessage::Notification(notification) => { - let params = match notification.params.clone() { + Ok(result) + } + crate::schema::v1::Response::Error { error, .. } => Err(error), + }; + } + RawJsonRpcMessage::Notification(notification) => { + check_payload_size(¬ification, MAX_PAYLOAD_BYTES)?; + let params = + match notification.params { Some(params) => match params.into_value() { Value::Object(map) => Some(map), _ => return Err(crate::Error::invalid_params().data( @@ -258,25 +306,20 @@ where }, None => None, }; - connection_for_task.send_notification_to( - Agent, - Protocol::notification( - server_id.clone(), - request_id.clone(), - notification.method.to_string(), - params, - ), - )?; - } - RawJsonRpcMessage::Request(_) => { - return Err(crate::Error::method_not_found() - .data("reverse MCP requests are not supported")); - } + connection_for_task.send_notification_to( + Agent, + Protocol::notification( + server_id.clone(), + request_id.clone(), + notification.method.to_string(), + params, + ), + )?; + } + RawJsonRpcMessage::Request(_) => { + return Err(crate::Error::method_not_found() + .data("reverse MCP requests are not supported")); } - Ok(()) - })?; - if let Some(response) = result { - return response; } } Err(crate::util::internal_error( @@ -312,10 +355,8 @@ where } Ok(()) }); - if let Err(error) = spawn_result { - // The dropped task also drops its responder and backend stop sender. - return Err(error); - } + // A failed spawn drops its responder and backend stop sender with the task. + spawn_result?; Ok(Handled::Yes) } } @@ -346,34 +387,68 @@ where } } +/// Discovery describes the revisions available through this binding, not other +/// transports the hosted backend might also implement. +fn constrain_discovery_versions(result: &mut Value) -> Result<(), crate::Error> { + let versions = result + .get_mut("supportedVersions") + .and_then(Value::as_array_mut) + .ok_or_else(|| crate::Error::internal_error().data("invalid MCP discovery result"))?; + if !versions + .iter() + .any(|version| version.as_str() == Some(MCP_VERSION)) + { + return Err(crate::Error::new(-32022, "Unsupported protocol version") + .data(serde_json::json!({"requested": MCP_VERSION, "supported": versions}))); + } + *versions = vec![Value::String(MCP_VERSION.to_owned())]; + Ok(()) +} + fn validate_modern_request( method: &str, params: Option<&Map>, ) -> Result<(), crate::Error> { if method == "initialize" { return Err( - crate::Error::invalid_params().data("native MCP requests do not use initialize") + crate::Error::method_not_found().data("native MCP requests do not use initialize") ); } let meta = params .and_then(|params| params.get("_meta")) - .and_then(Value::as_object); - if meta - .and_then(|meta| meta.get("io.modelcontextprotocol/protocolVersion")) + .and_then(Value::as_object) + .ok_or_else(|| { + crate::Error::invalid_params().data("inner params._meta must be an object") + })?; + let version = meta + .get("io.modelcontextprotocol/protocolVersion") .and_then(Value::as_str) - != Some("2026-07-28") - || !meta - .and_then(|meta| meta.get("io.modelcontextprotocol/clientCapabilities")) - .is_some_and(Value::is_object) + .ok_or_else(|| { + crate::Error::invalid_params() + .data("inner params._meta requires io.modelcontextprotocol/protocolVersion") + })?; + if version != MCP_VERSION { + return Err(crate::Error::new(-32022, "Unsupported protocol version") + .data(serde_json::json!({"requested": version, "supported": [MCP_VERSION]}))); + } + if !meta + .get("io.modelcontextprotocol/clientCapabilities") + .is_some_and(Value::is_object) { - return Err(crate::Error::invalid_params().data("inner params._meta requires io.modelcontextprotocol/protocolVersion 2026-07-28 and io.modelcontextprotocol/clientCapabilities object")); + return Err(crate::Error::invalid_params().data( + "inner params._meta requires io.modelcontextprotocol/clientCapabilities object", + )); } Ok(()) } #[cfg(test)] mod tests { - use super::validate_modern_request; + use super::{ + ActiveRequests, MAX_ACTIVE_REQUESTS, admit_request, check_payload_size, + constrain_discovery_versions, validate_modern_request, + }; + use crate::schema::v1::McpRequestId; use serde_json::json; #[test] @@ -391,4 +466,73 @@ mod tests { .is_err() ); } + + #[test] + fn unsupported_version_is_an_mcp_error_not_a_legacy_fallback() { + let params = json!({ + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2025-11-25", + "io.modelcontextprotocol/clientCapabilities": {} + } + }); + let error = validate_modern_request("tools/call", params.as_object()).unwrap_err(); + assert_eq!( + serde_json::to_value(error).unwrap(), + json!({ + "code": -32022, + "message": "Unsupported protocol version", + "data": {"requested": "2025-11-25", "supported": ["2026-07-28"]} + }) + ); + } + + #[test] + fn native_request_admission_is_bounded_and_recovers_after_cleanup() { + let active = ActiveRequests::default(); + let mut admitted = Vec::new(); + for index in 0..MAX_ACTIVE_REQUESTS { + admitted + .push(admit_request(&active, McpRequestId::new(format!("req-{index}"))).unwrap()); + } + assert_eq!(active.lock().unwrap().len(), MAX_ACTIVE_REQUESTS); + let duplicate = admit_request(&active, McpRequestId::new("req-0")) + .err() + .unwrap(); + assert_eq!(duplicate.code, crate::ErrorCode::InvalidParams); + let overload = admit_request(&active, McpRequestId::new("extra")) + .err() + .unwrap(); + assert_eq!(i32::from(overload.code), -32000); + drop(admitted.pop()); + let replacement = admit_request(&active, McpRequestId::new("replacement")).unwrap(); + assert_eq!(active.lock().unwrap().len(), MAX_ACTIVE_REQUESTS); + drop(replacement); + drop(admitted); + assert!(active.lock().unwrap().is_empty()); + } + + #[test] + fn payload_limits_count_json_escaping_without_building_an_extra_buffer() { + let payload = json!({"text": "\n\n"}); + let encoded = serde_json::to_vec(&payload).unwrap(); + assert!(check_payload_size(&payload, encoded.len()).is_ok()); + assert!(check_payload_size(&payload, encoded.len() - 1).is_err()); + } + + #[test] + fn discovery_reports_the_binding_version_without_changing_other_payload() { + let mut result = json!({ + "resultType": "complete", + "supportedVersions": ["2025-11-25", "2026-07-28"], + "capabilities": {"tools": {}}, + "_meta": {"vendor/opaque": ["preserved"]} + }); + constrain_discovery_versions(&mut result).unwrap(); + assert_eq!(result["supportedVersions"], json!(["2026-07-28"])); + assert_eq!(result["_meta"]["vendor/opaque"], json!(["preserved"])); + assert_eq!(result["capabilities"], json!({"tools": {}})); + let mut unsupported = json!({"supportedVersions": ["2025-11-25"]}); + assert!(constrain_discovery_versions(&mut unsupported).is_err()); + assert!(constrain_discovery_versions(&mut json!({})).is_err()); + } } From 85571cc9ccacfd05b7e6f512be3231e0bb4c030f Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Fri, 25 Sep 2026 11:45:03 +0200 Subject: [PATCH 3/8] feat(acp)!: own MCP request lifetimes and bound transport queues Separate MCP outcomes from ACP failures, reuse request-native services, join tool cleanup before releasing admission, and bound retained frames, replies and HTTP response bodies. Preserve passive transport half-close semantics and test real rmcp HTTP workflows. BREAKING CHANGE: Channel now carries budgeted frames, ConnectTo returns ConnectionDriver, and native MCP uses request-scoped services with explicit outcome carriers. Coordinate the schema and dependent SDK major releases before publishing. --- Cargo.lock | 26 +- Cargo.toml | 3 +- md/SUMMARY.md | 1 + md/mcp-bridge.md | 82 +- md/mcp-over-acp.md | 85 +- md/migration-rmcp-v4.md | 11 +- md/migration-stateless-mcp.md | 133 + md/protocol.md | 42 +- md/transport-architecture.md | 25 +- .../src/trace.rs | 91 +- .../tests/mcp_cleanup_ownership.rs | 294 ++ .../tests/mcp_over_acp_polyfill.rs | 13 +- .../tests/mcp_over_acp_polyfill_v2.rs | 48 +- .../tests/stateless_mcp_http.rs | 311 +++ .../tests/test_tool_fn.rs | 142 + .../tests/trace_client_mcp_server.rs | 3 + .../tests/trace_mcp_tool_call.rs | 82 +- .../tests/trace_snapshot.rs | 3 + src/agent-client-protocol-http/src/client.rs | 520 +++- .../src/client_admission_tests.rs | 297 ++ .../src/connection.rs | 394 ++- .../src/connection_admission_tests.rs | 101 + .../src/http_server.rs | 86 +- .../src/protocol.rs | 15 +- .../src/websocket_server.rs | 41 +- src/agent-client-protocol-polyfill/Cargo.toml | 2 + .../src/mcp_over_acp/http.rs | 480 +++- .../src/mcp_over_acp/mod.rs | 334 +-- .../examples/stateless_native_mcp.rs | 8 +- src/agent-client-protocol-rmcp/src/builder.rs | 40 +- src/agent-client-protocol-rmcp/src/lib.rs | 71 +- src/agent-client-protocol-rmcp/src/native.rs | 239 ++ .../tests/stateless_native_mcp.rs | 157 +- src/agent-client-protocol/Cargo.toml | 1 + .../examples/v2_session_coordination/tests.rs | 24 +- src/agent-client-protocol/src/acp_agent.rs | 4 +- src/agent-client-protocol/src/component.rs | 71 +- src/agent-client-protocol/src/jsonrpc.rs | 2389 +++++++++++++++-- .../src/jsonrpc/admission.rs | 274 ++ .../src/jsonrpc/incoming_actor.rs | 87 +- .../src/jsonrpc/outgoing_actor.rs | 199 +- .../src/jsonrpc/task_actor.rs | 19 +- .../src/jsonrpc/transport_actor.rs | 148 +- src/agent-client-protocol/src/lib.rs | 15 +- .../src/mcp_server/active_session.rs | 410 ++- .../src/mcp_server/context.rs | 23 + .../src/mcp_server/mod.rs | 13 + .../src/mcp_server/server.rs | 104 +- .../src/mcp_server/service.rs | 197 ++ .../src/mcp_server/tool_fn.rs | 116 +- src/agent-client-protocol/src/role/acp.rs | 23 +- .../src/schema/v2_impls.rs | 22 +- src/agent-client-protocol/src/util.rs | 74 - .../tests/application_dispatch_v2.rs | 18 +- .../tests/jsonrpc_advanced.rs | 7 +- .../tests/jsonrpc_batch.rs | 5 +- .../tests/jsonrpc_error_handling.rs | 7 +- .../tests/jsonrpc_transport_close.rs | 40 +- .../tests/protocol_v2.rs | 31 +- .../tests/proxy_protocol_router_v2.rs | 6 +- .../tests/session_ordering.rs | 50 +- .../tests/session_restore.rs | 50 +- .../tests/session_v2_mcp.rs | 8 +- 63 files changed, 7310 insertions(+), 1305 deletions(-) create mode 100644 md/migration-stateless-mcp.md create mode 100644 src/agent-client-protocol-conductor/tests/mcp_cleanup_ownership.rs create mode 100644 src/agent-client-protocol-conductor/tests/stateless_mcp_http.rs create mode 100644 src/agent-client-protocol-http/src/client_admission_tests.rs create mode 100644 src/agent-client-protocol-http/src/connection_admission_tests.rs create mode 100644 src/agent-client-protocol-rmcp/src/native.rs create mode 100644 src/agent-client-protocol/src/jsonrpc/admission.rs create mode 100644 src/agent-client-protocol/src/mcp_server/service.rs diff --git a/Cargo.lock b/Cargo.lock index 3bce6c5f..ed05ae40 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -9,6 +9,7 @@ dependencies = [ "agent-client-protocol-derive", "agent-client-protocol-schema", "agent-client-protocol-test", + "async-channel", "async-io", "async-process", "blocking", @@ -112,7 +113,9 @@ dependencies = [ "axum", "base64 0.23.1", "futures", + "hmac", "serde_json", + "sha2", "tokio", "tracing", "uuid", @@ -138,7 +141,7 @@ dependencies = [ [[package]] name = "agent-client-protocol-schema" version = "1.9.1" -source = "git+https://github.com/agentclientprotocol/agent-client-protocol?rev=1ae7f09519fa0ba43289365da42bd589468e135f#1ae7f09519fa0ba43289365da42bd589468e135f" +source = "git+https://github.com/agentclientprotocol/agent-client-protocol?rev=9d8b499a332ffa0655007d4be7dd949e05180de3#9d8b499a332ffa0655007d4be7dd949e05180de3" dependencies = [ "anyhow", "derive_more", @@ -875,6 +878,7 @@ checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ "block-buffer 0.10.4", "crypto-common 0.1.7", + "subtle", ] [[package]] @@ -1200,6 +1204,15 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" +[[package]] +name = "hmac" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" +dependencies = [ + "digest 0.10.7", +] + [[package]] name = "http" version = "1.5.0" @@ -2533,6 +2546,17 @@ dependencies = [ "digest 0.11.3", ] +[[package]] +name = "sha2" +version = "0.10.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest 0.10.7", +] + [[package]] name = "sharded-slab" version = "0.1.7" diff --git a/Cargo.toml b/Cargo.toml index 141e7efc..efa72e48 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -36,7 +36,7 @@ yopo = { package = "agent-client-protocol-yopo", path = "src/yopo" } # Protocol # Draft cross-repository validation; replace with the released schema before publishing. -agent-client-protocol-schema = { git = "https://github.com/agentclientprotocol/agent-client-protocol", rev = "1ae7f09519fa0ba43289365da42bd589468e135f", default-features = false, features = ["tracing"] } +agent-client-protocol-schema = { git = "https://github.com/agentclientprotocol/agent-client-protocol", rev = "9d8b499a332ffa0655007d4be7dd949e05180de3", default-features = false, features = ["tracing"] } # Core async runtime tokio = { version = "1.52", default-features = false } @@ -73,6 +73,7 @@ url = "2.5" async-io = "2" async-process = "2" async-stream = "0.3.6" +async-channel = "2" blocking = "1" chrono = "0.4" futures = "0.3.32" diff --git a/md/SUMMARY.md b/md/SUMMARY.md index 0e619d4a..31ed83c7 100644 --- a/md/SUMMARY.md +++ b/md/SUMMARY.md @@ -32,6 +32,7 @@ # Reference +- [Migrating the Native MCP Transport](./migration-stateless-mcp.md) - [Migrating the rmcp Integration to v4](./migration-rmcp-v4.md) - [Migrating to v2.0](./migration_v2.0.md) - [Migrating to v0.11](./migration_v0.11.x.md) diff --git a/md/mcp-bridge.md b/md/mcp-bridge.md index 74fa9f48..af412d3c 100644 --- a/md/mcp-bridge.md +++ b/md/mcp-bridge.md @@ -92,14 +92,15 @@ support and rejects any native declaration that is nevertheless supplied. For each schema-selected `McpServer::Acp` entry in a session setup request, the polyfill: -1. Creates or reuses a connection-scoped localhost bridge endpoint for the - `serverId` and replaces the declaration with the HTTP transport for the - final agent. -2. Retains the native `serverId` so connections can be routed back to the - component that provided the server. -3. Adds a runtime-only bearer credential to the HTTP declaration. The endpoint - requires that credential and checks supplied Origin headers; an ephemeral - port alone is not access control. +1. Creates or reuses one connection-scoped loopback listener and replaces the + declaration with an HTTP URL whose path encodes the non-secret `serverId`. + No per-server listener or route-table entry is allocated. +2. Routes each request back to the component that owns that native registration. + The provider, not possession of the URL, decides whether it still exists. +3. Adds a runtime-only bearer credential derived from the connection secret and + server ID to the HTTP declaration's headers. The endpoint authenticates and + checks supplied Origin headers before reading the request body. Credentials + never appear in URLs; an ephemeral port alone is not access control. 4. For each POST, allocates a unique logical MCP request ID and sends `mcp/message` to the provider. Two HTTP clients may use the same external JSON-RPC ID without sharing routing or state. @@ -114,10 +115,11 @@ versions include `session/fork` when `unstable_session_fork` is enabled. Declarations using another transport are left unchanged, including extension transports represented by v2's `McpServer::Other`. -Endpoints are cached by `serverId` across session setup requests on the ACP -connection. The output declaration is rebuilt for each occurrence, preserving -that occurrence's `name`, `_meta`, and other unmodified extension fields even -when its endpoint is reused. +The same server ID derives the same route and credential on this ACP connection. +The output declaration is rebuilt for each occurrence, preserving its `name`, +`_meta`, and other unmodified extension fields. Failed setup and declaration +churn cannot accumulate per-server endpoint allocations. A server ID must never +be rebound to a different registration during the connection's lifetime. The native wire envelopes are documented in the [SDK Protocol Reference](./protocol.md#native-mcp-over-acp). @@ -125,8 +127,8 @@ Reference](./protocol.md#native-mcp-over-acp). ## HTTP Mode `McpOverAcpPolyfill::http()` is the default compatibility shape. It replaces -the native declaration with an HTTP MCP URL at `http://127.0.0.1:PORT`. The -embedded server accepts a single JSON-RPC request per POST at `/`, returning +the native declaration with an HTTP MCP URL at `http://127.0.0.1:PORT/`. The +embedded server accepts a single JSON-RPC request per POST at that route, returning JSON for a terminal-only response or SSE for a request that emits notifications. GET and DELETE return 405. Batches and client-originated JSON-RPC responses are rejected; there is no standalone GET event stream or MCP session ID. @@ -148,29 +150,47 @@ other metadata, progress tokens, and opaque retry state are not rewritten. ## Lifecycle and Failure Behavior -Each POST owns a pending native request, not an MCP session. A terminal result, -error, response-stream close, or overflow removes that request's routing state. -The listening endpoint remains available for later requests. +Each POST owns a pending native request, not an MCP session. Closing its response +stream cancels that request. A terminal outcome ends native work, but HTTP +admission remains held until the response body is consumed or dropped. The +listening endpoint remains available for later requests; releasing the native +registration makes requests through its old URL fail rather than reviving it. The adapter limits each response's queued notifications to 16 messages and -256 KiB of serialized data, with 64 active requests and 32 listening endpoints -per adapter. A separate terminal-response path avoids stranding completion -behind a full queue. Overflow explicitly fails and cancels that operation -without blocking the shared runner or dropping events silently. +256 KiB of serialized data, admits at most 64 HTTP responses at a time, and caps +request bodies and terminal payloads at 1 MiB. The body owns the admission permit, +including while a client is not reading. A separate terminal-response path +avoids stranding completion behind a full queue. Overflow explicitly fails and +cancels that operation without blocking the shared runner or dropping events silently. + +The bridge unwraps the ACP outcome carrier before creating the HTTP JSON-RPC +response. MCP error codes/data stay MCP errors; binding failures use their +separate error codes. Queued notifications precede the terminal response. Unknown or late provider notifications are ignored; reverse MCP requests are not supported. The adapter does not infer ACP session IDs or maintain MCP initialization state. -## Remaining scope +## Native-tool re-export contract + +The adapter creates a **new HTTP endpoint for native tool semantics**. It does +not preserve another HTTP gateway's parameter-header routing or authorization. +It removes transport-only `x-mcp-header` annotations from actual schema positions +in `tools/list` results. Argument schemas and validation keywords, tool ordering, +pagination, metadata, and similarly named properties/example/default data remain +unchanged. Annotated native tools remain listed and callable. + +Each `tools/call` issues exactly one native call, without hidden descriptor reads +or a prior client `tools/list` requirement. Native passthrough does not transform +the original descriptors. `Mcp-Param-*` headers are rejected; they confer no +authority on this endpoint. Standard MCP method/name/version header checks remain. + +If a deployment depends on an existing HTTP gateway's mirrored-parameter policy, +it must implement that policy at this endpoint or decline this re-export. -Tools using `x-mcp-header` annotations are currently unsupported and fail -closed: they are omitted from listings, calls are rejected, and supplied -`Mcp-Param-*` headers are rejected. For a direct tool call the adapter fetches -the tool descriptor internally, including pagination, so the caller does not -need a prior tools/list handshake. That lookup is an explicit per-call cost. +## Validation scope -This is not yet full HTTP conformance. Native SDK `Channel` and outgoing -queues also remain unbounded; the HTTP queue limits above do not establish -end-to-end native backpressure. Native admission/payload limits and the -remaining transport work are described in [Native MCP-over-ACP](./mcp-over-acp.md). +This does not establish every optional MCP feature or complete HTTP conformance. +In particular, HTTP response limits alone do not prove native transport bounds. +Owned operation cleanup and end-to-end bounded transport are stabilization gates; +see [Native MCP-over-ACP](./mcp-over-acp.md). diff --git a/md/mcp-over-acp.md b/md/mcp-over-acp.md index 190ea7ae..09e479da 100644 --- a/md/mcp-over-acp.md +++ b/md/mcp-over-acp.md @@ -15,14 +15,31 @@ Attach an `mcp_server::McpServer` to session setup through the existing builder APIs. It publishes a `McpServer::Acp` declaration with a provider-generated `serverId`. -Each incoming `mcp/message` invokes the backend factory for one operation. -The MCP request context exposes `server_id()` and `request_id()`; standalone -MCP serving has neither. Tool definitions can be shared, but per-request MCP -metadata and capabilities must not be inferred from previous operations. - -The rmcp integration can construct tools through its builder or wrap a supplied -rmcp 3.4 service. The normal rmcp service can process a modern request without -`initialize` when its inner `_meta` declares the modern version and capabilities. +`McpService` is a reusable application service. Each `execute` call owns one +operation future and receives an `McpRequestContext` with `server_id()`, +`request_id()`, validated `metadata()`, cancellation, and an async +`send_notification` method. Share tool implementations, caches, and connection +pools deliberately; never infer a request's identity or capabilities from a +previous operation. + +Use `McpServer::new_service` for a native service, or +`new_service_with_standalone` when also exposing an independent standalone +transport. The connector-based factory remains an explicit adapter for backends +that require per-operation construction; stateless MCP does not require it. + +The rmcp integration's builder and `from_rmcp` use the reusable service path +for ACP attachments. Each operation uses rmcp's direct, one-request transport +without `initialize`. Its wrapper supervises rmcp handler futures through +cancellation and cleanup instead of merely dropping detached task handles. + +Custom `McpService` implementations must observe `operation_cancellation()` and +return only after their owned cleanup finishes. The binding waits for this +completion; it cannot forcibly terminate detached application work. + +The scoped `tool_fn` helpers continue to provide `McpConnectionTo` for host ACP +access. For decisions using the full MCP metadata/capabilities, implement +`McpService` or an rmcp handler receiving its `RequestContext`. Standalone MCP +connections have no ACP server or logical request ID. ## Consuming tools @@ -44,9 +61,17 @@ stream notifications. Route by server and logical request ID. Do not block the ACP dispatch loop waiting for peer traffic; use a spawned task or the connection's application future. -The final response is the MCP result directly, including its `resultType`, or -the original MCP error. For MRTR, process the `input_required` result and send -a fresh request with `inputResponses` and the exact opaque `requestState`. +The final successful ACP response is `MessageMcpResponse::Result { result, .. }` +or `MessageMcpResponse::Error { error, .. }`. Match that carrier before interpreting +the MCP outcome. The result preserves all MCP fields, including `resultType`; +the error preserves its MCP code, message, optional data, and extensions. +An MCP code must never be treated as an ACP code: for example, inner `-32000` +does not mean ACP authentication is required. + +Outer ACP failures instead describe invalid binding input, cancellation, +resource exhaustion, an unavailable registration, or a failed backend/transport. +For MRTR, process the inner `input_required` result and send a fresh request +with `inputResponses` and the exact opaque `requestState`. Discovery reports only the MCP revision exposed by this binding, even if the hosted backend also supports older revisions through other transports. @@ -59,22 +84,35 @@ arrive as request-scoped notifications, with the logical request ID in that subscription's state or lifetime. Use `SentRequest::cancel` (or drop an unconsumed request) to cancel the outer -ACP operation. The provider stops that operation's backend work and returns a -result or cancellation error. Removing a provider stops its outstanding work; -no separate `mcp/disconnect` exchange exists. +ACP operation. The provider revokes output immediately and stops that operation's +owned backend work; its admission slot and logical ID remain held until cleanup +finishes. Cancellation produces an outer cancellation error unless completion +already won the race. Removing a registration or receiving transport EOF cancels +its outstanding work; no separate `mcp/disconnect` exchange exists. ## Resource limits and remaining work -The native provider admits at most 64 concurrent operations per declared -server and checks a 16 MiB serialized payload limit before starting work or -forwarding backend responses/notifications. Rejected work reports an error; -completion and cancellation release the admission slot. +The native binding has per-registration admission and serialized payload limits. +Resource exhaustion is an outer `MCP_RESOURCE_EXHAUSTED` (`-33000`) failure, not +ACP authentication and not an inner MCP tool error. + +The transport revision introduces finite `ConnectionLimits` and `BudgetedFrame` +ownership. Adapters must keep the frame's permit through staging, deferred +dispatch, and writes; extracting a payload must not silently release its charge +while retaining the data. Async producers await capacity; synchronous dispatch +must fail explicitly instead of blocking the dispatcher needed to free capacity. + +The same item-limit policy currently governs frame queues, pending requests, +running tasks, dynamic handlers, and deferred dispatch; the default is 32. +The shared payload budget defaults to 64 MiB with a 16 MiB frame maximum and +reserved response/cancellation capacity. These are serialized-payload charges, +not an exact bound on total process memory or allocations inside user code. -These are not end-to-end memory bounds. The public SDK `Channel` and outgoing -queues remain unbounded. A bounded native transport path is still required -before stabilization; admission and per-message size checks do not prevent -accumulation behind a slow peer. The [HTTP adapter](./mcp-bridge.md) separately -bounds its own response queues and fails/cancels an overflowing operation. +Regression coverage includes sender-clone saturation, cross-budget forwarding, +retained responses and callbacks, EOF draining, and cancellation while cleanup is +paused. The [HTTP adapter](./mcp-bridge.md) separately owns its response-body permits +and fails/cancels overflowing operations. Full MCP conformance and protocol +stabilization remain separate from this implementation evidence. ## Runnable example @@ -87,3 +125,4 @@ cargo run -p agent-client-protocol-rmcp \ This direct ACP example uses actual rmcp tools without the HTTP polyfill. See the [protocol reference](./protocol.md#native-mcp-over-acp) for wire details and the [RFD](https://agentclientprotocol.com/rfds/mcp-over-acp) for the design. +The [migration guide](./migration-stateless-mcp.md) lists the breaking changes. diff --git a/md/migration-rmcp-v4.md b/md/migration-rmcp-v4.md index 1de6f435..d0abc4ca 100644 --- a/md/migration-rmcp-v4.md +++ b/md/migration-rmcp-v4.md @@ -54,8 +54,9 @@ the required request metadata. ## MCP-over-ACP remains a separate draft -This upgrade does not change ACP's unstable `mcp/connect`, `mcp/message`, or -`mcp/disconnect` envelopes. Redesigning that transport around stateless, -server-addressed requests is separate work. The new transport's latest-only -target does not require removing existing rmcp behavior from this prerequisite -dependency upgrade. +The dependency upgrade alone did not change the unstable ACP wire envelopes. +The subsequent [native MCP transport migration](./migration-stateless-mcp.md) +removes the prototype's `mcp/connect`/`mcp/disconnect` lifecycle and changes +`mcp/message` to server-addressed requests with explicit outcome carriers. +Read both guides when adopting the combined major-version changes. The +latest-only native binding does not require removing standalone rmcp behavior. diff --git a/md/migration-stateless-mcp.md b/md/migration-stateless-mcp.md new file mode 100644 index 00000000..6b26975f --- /dev/null +++ b/md/migration-stateless-mcp.md @@ -0,0 +1,133 @@ +# Migrating the Native MCP Transport + +This draft replaces the connection-oriented MCP-over-ACP prototype with a +request-scoped binding for **MCP 2026-07-28 only**. It is part of the next major +SDK change, not a compatibility layer for older MCP revisions. The +`unstable_mcp_over_acp` gate remains; draft ACP v2 still has its separate gate. + +## Wire changes + +| Previous prototype | New binding | +| --- | --- | +| `mcp/connect` and `mcp/disconnect` | Removed | +| `McpConnectionId` / `connectionId` | Removed | +| `mcp/message(connectionId, method, params)` | `mcp/message(serverId, requestId, method, params)` | +| MCP initialization and connection-scoped capabilities | Required version/capabilities in each request's inner `_meta` | +| Arbitrary reverse MCP requests | MRTR `input_required` results and explicit caller retries | +| Raw MCP result or MCP error in the ACP error envelope | Successful ACP response containing exactly one inner `result` or `error` | +| HTTP MCP sessions and standalone GET streams | Independent POSTs, including long-lived subscription POSTs | + +Keep the server declaration's `serverId`. Generate a fresh logical +`McpRequestId` per call and pass it to +`MessageMcpRequest::new(server_id, request_id, method)`. That ID stays unchanged +through proxies; it is not the hop-local ACP JSON-RPC ID. + +Every inner request includes: + +```json +{ + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } +} +``` + +There is no hidden initialization or discovery prerequisite. Explicitly select +2026-07-28 when constructing an rmcp client: rmcp 3.4's default version constant +still selects an older revision. + +## Handle two error domains + +First handle the outer ACP request result, then match `MessageMcpResponse`: + +- `Result { result, .. }` contains an opaque MCP result, including any MCP + metadata, `resultType`, or explicit JSON null. +- `Error { error, .. }` contains an `McpError`. Its `code` is a plain MCP integer, + not ACP's `ErrorCode`. `data` preserves omission separately from JSON null, + and unknown error extensions survive. +- An outer ACP error reports a binding failure: invalid envelope, cancellation, + resource limit, unavailable registration, or backend/transport failure. + +Do not run ACP authentication handling on an inner MCP error code. A tool +execution failure with `isError` remains an MCP result. MRTR's `input_required` +also remains a result; retry with fresh IDs/metadata and unchanged opaque state. + +Both ACP versions export the same response/error carrier types. Downstream +code that implements traits for these types must not provide separate v1 and +v2 implementations. + +## Separate services from operations + +Use the reusable `McpService` abstraction for native providers. Per-operation +`McpRequestContext` contains logical/server identity, MCP metadata/capabilities, +cancellation, and request-scoped notification permissions. A service can share +application state without sharing MCP protocol state. + +`McpServer::new_service` registers a native service. An explicit factory/standalone +adapter remains available when constructing a backend per operation is actually +needed. `McpServer::from_rmcp` and the rmcp tool builder retain their attachment +entry points but execute ACP requests through the request-native service path. + +Do not detach tool work from its operation. Cancelling a queued call must prevent +it from starting; cancelling a running call must drop or stop its owned future +and supervise cleanup. Failure to deliver a cancelled tool's result must not +terminate the containing ACP connection. + +## Registration and cancellation + +A server ID names one registration during an ACP connection's lifetime. Do not +rebind a removed ID to another provider. Dropping the local registration rejects +future calls and cancels its active work; omitting a declaration from a later +setup request is not a new unadvertisement message. + +Use ACP request cancellation, not an MCP disconnect. Cancellation revokes +notifications immediately while cleanup retains the active ID and admission +permit. Independent calls, subscriptions, and the reusable service remain alive. +Transport EOF must begin this cleanup even if application code is still waiting +on the disconnected peer. + +## HTTP clients + +The local polyfill re-exports native tools through one signed, loopback HTTP +endpoint per ACP connection. Pass the declaration's Authorization header, never +put its bearer credential in a URL. Requests use current MCP headers and do not +exchange session IDs or `initialize`. + +The endpoint strips transport-only `x-mcp-header` schema annotations from tool +descriptors and rejects `Mcp-Param-*` headers. It does not transport an existing +HTTP gateway's routing or authorization policy. Direct tool calls require no +preliminary descriptor fetch. See the [HTTP adapter contract](./mcp-bridge.md). + +## Custom transports and connectors + +`ConnectTo::into_channel_and_future` now returns `(Channel, ConnectionDriver)`. +Wrap an owned driver future in `ConnectionDriver::new`; use +`ConnectionDriver::passive()` only for an endpoint driven elsewhere. Awaiting +the driver remains supported. Do not treat passive-driver completion as EOF: +doing so drops final responses when an input stream half-closes. + +`Channel::rx` yields `BudgetedFrame`, not a bare wire frame. Use `.frame()` to +inspect it, and preserve the envelope when forwarding through a sink. If a +custom adapter separates payload from accounting with `.into_parts()`, retain +the permit as long as the deferred payload or serialized output exists. + +For a new raw frame, use `FrameSender::send_frame(frame).await` outside dispatch +or `try_send(frame)` for explicit fail-fast admission. The old `unbounded_send` +API is removed; ignoring capacity errors silently loses protocol traffic. +Finite queue and byte policies are configured through `ConnectionLimits`. + +## Release checklist + +- Replace the draft Git schema pin with the released matching schema version. +- Coordinate major releases for crates whose public transport or rmcp-facing + API changed; do not infer compatibility solely from unchanged Cargo numbers. +- Exercise v1 and v2 carrier/error behavior, cancellation and EOF, MRTR, + subscriptions, and slow consumers before stabilizing. +- Follow the bounded transport API's ownership rules when writing custom + adapters: moving a payload must not release its accounting while a deferred + dispatch, writer, or unread HTTP body still retains it. + +The [native guide](./mcp-over-acp.md) and [protocol reference](./protocol.md#native-mcp-over-acp) +describe the target behavior. Historical migration chapters describe earlier +releases and are not a specification for this binding. diff --git a/md/protocol.md b/md/protocol.md index 8b5751cd..9ff0a2f3 100644 --- a/md/protocol.md +++ b/md/protocol.md @@ -74,8 +74,8 @@ opt-in `session/fork`): ``` `serverId` identifies the declared server and is used to route `mcp/message` -back to the component that provided it. A provider must not reuse one server ID -for multiple visible servers on the same ACP connection. The high-level +back to the component that provided it. A provider must not rebind a server ID +to another registration on the same ACP connection, even after removal. The high-level `agent_client_protocol::mcp_server::McpServer` APIs create this declaration automatically. @@ -112,9 +112,30 @@ JSON-RPC ID is renumbered: } ``` -The outer response carries the inner MCP result (including `resultType`) or -error directly. MRTR `input_required` is a result, not a reverse RPC; retry the -original operation with fresh metadata/IDs and unchanged opaque state. +The successful outer ACP response contains exactly one MCP outcome: + +```json +{ + "jsonrpc": "2.0", + "id": 21, + "result": { + "result": { "resultType": "complete", "content": [] } + } +} +``` + +An MCP protocol error uses `{"error": {"code": ..., "message": ..., "data": ...}}` +inside the successful outer `result`, not an ACP error response. The shared +`MessageMcpResponse::{Result, Error}` type preserves this distinction. Inner +results are opaque JSON (including null); inner error data distinguishes null +from omission. MCP error codes never acquire ACP meanings. + +Outer ACP errors describe binding failures: invalid envelope/duplicate ID +(`-32602`), cancellation (`-32800`), resource exhaustion (`-33000`), unavailable +registration (`-33001`), or backend/transport failure (`-33002`). + +MRTR `input_required` is an MCP result, not a reverse RPC; retry the original +operation with fresh metadata/IDs and unchanged opaque state. For `server/discover`, supported versions are restricted to the revision exposed by this binding; a backend must actually support that revision. @@ -148,12 +169,15 @@ mean no parameters. A valid modern request still needs its required Use [`$/cancel_request`](./request-cancellation.md) with the outer ACP request ID. Normal proxy forwarding maps this cancellation hop by hop. It never -rewrites the logical MCP ID. +rewrites the logical MCP ID. Advertising this binding requires cancellation +handling even where the underlying ACP revision makes general cancellation optional. Each operation owns its backend work. A result, error, cancellation, or -provider removal ends that operation; sibling requests and subscriptions stay -independent. There is no MCP connection ID to release. `server/discover` is an -ordinary optional request, not a prerequisite for tool calls. +registration removal ends that operation; sibling requests and subscriptions stay +independent. Cancellation revokes output immediately, but the operation keeps its +admission slot and logical ID until owned cleanup finishes. There is no MCP +connection ID to release. `server/discover` is an ordinary optional request, +not a prerequisite for tool calls. ## Related Documentation diff --git a/md/transport-architecture.md b/md/transport-architecture.md index 8204b77e..ef9b84fe 100644 --- a/md/transport-architecture.md +++ b/md/transport-architecture.md @@ -262,16 +262,29 @@ Ordering](./conductor.md#routing-and-ordering). is the common component and transport abstraction. `connect_to` joins a component to its counterpart and drives the connection until completion. `into_channel_and_future` exposes the canonical low-level boundary as a -`Channel` plus the future that drives the component: +`Channel` plus an explicit connection driver: ```rust,ignore -fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<()>>); +fn into_channel_and_future(self) -> (Channel, ConnectionDriver); ``` -The returned future owns transport failures and lifecycle completion. The -channel carries only `TransportFrame` wire events. Most components implement -only `connect_to`; direct transports override `into_channel_and_future` to avoid -an intermediate copy. +`ConnectionDriver::new(future)` owns transport failures and component +completion. `ConnectionDriver::passive()` denotes an endpoint driven elsewhere, +such as an existing `Channel`; its no-op completion is **not** an EOF signal. +Both implement `Future`, so drivers can still be joined with application work. +Dynamic connectors preserve this distinction. + +A bridge must poll both copy directions while an active component runs. When +the component finishes, drain its accepted output without requiring the remote +sender to close. Between two passive endpoints, preserve independent half-close: +input EOF must still allow a final response in the other direction. + +The channel carries `BudgetedFrame` values containing complete `TransportFrame` +wire events and their resource permits. Most components implement only +`connect_to`; direct transports override `into_channel_and_future` to avoid +an intermediate copy. A forwarded frame keeps its permit through any adapter +queue, deferred dispatch, or writer. This accounting is internal and does not +change the JSON-RPC wire shape. ## Transport Implementations diff --git a/src/agent-client-protocol-conductor/src/trace.rs b/src/agent-client-protocol-conductor/src/trace.rs index 3921e954..e83fa469 100644 --- a/src/agent-client-protocol-conductor/src/trace.rs +++ b/src/agent-client-protocol-conductor/src/trace.rs @@ -11,7 +11,7 @@ use std::time::Instant; use agent_client_protocol::schema::SuccessorMessage; use agent_client_protocol::schema::v1::{ - MessageMcpNotification, MessageMcpRequest, Notification as RpcNotification, + MessageMcpNotification, MessageMcpRequest, MessageMcpResponse, Notification as RpcNotification, Request as RpcRequest, RequestId, Response as RpcResponse, }; use agent_client_protocol::{ @@ -98,6 +98,11 @@ pub struct ResponseEvent { /// True if this is an error response. pub is_error: bool, + /// Whether an error belongs to the outer ACP binding or the inner MCP peer. + /// Older trace files omit this provenance. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub error_domain: Option, + /// Response result or error object. pub payload: serde_json::Value, } @@ -181,12 +186,7 @@ impl std::fmt::Debug for TraceWriter { } struct RequestDetails { - #[expect(dead_code)] protocol: Protocol, - - #[expect(dead_code)] - method: String, - request_from: ComponentIndex, request_to: ComponentIndex, } @@ -239,7 +239,6 @@ impl TraceWriter { id.clone(), RequestDetails { protocol, - method: method.clone(), request_from: from, request_to: to, }, @@ -262,7 +261,7 @@ impl TraceWriter { from: ComponentIndex, to: ComponentIndex, id: serde_json::Value, - is_error: bool, + error_domain: Option, mut payload: serde_json::Value, ) { redact_http_credentials(&mut payload); @@ -271,7 +270,8 @@ impl TraceWriter { from: format!("{from:?}"), to: format!("{to:?}"), id, - is_error, + is_error: error_domain.is_some(), + error_domain, payload, })); } @@ -370,13 +370,13 @@ impl TraceWriter { }; let id = id_to_json(&id); if let Some(RequestDetails { - protocol: _, - method: _, + protocol, request_from, request_to, }) = self.request_details.remove(&id) { - self.response(request_to, request_from, id, is_error, payload); + let (error_domain, payload) = response_outcome(protocol, is_error, payload); + self.response(request_to, request_from, id, error_domain, payload); } } } @@ -529,6 +529,33 @@ fn params_from_transport(params: Option) -> serde_json::Value params.map_or(serde_json::Value::Null, RawJsonRpcParams::into_value) } +/// Project the logical MCP outcome without losing whether an error came from +/// the outer ACP binding. In particular, an inner -32000 is not ACP AuthRequired. +fn response_outcome( + protocol: Protocol, + outer_error: bool, + payload: serde_json::Value, +) -> (Option, serde_json::Value) { + if outer_error { + return (Some(Protocol::Acp), payload); + } + if protocol == Protocol::Mcp { + match serde_json::from_value::(payload.clone()) { + Ok(MessageMcpResponse::Result { result, .. }) => return (None, result), + Ok(MessageMcpResponse::Error { error, .. }) => { + return ( + Some(Protocol::Mcp), + serde_json::to_value(error).expect("MCP errors contain only JSON values"), + ); + } + // Retain a malformed carrier as observed, rather than invent an + // error the peer never sent. The binding validates it separately. + _ => {} + } + } + (None, payload) +} + /// Do not persist HTTP credentials from MCP declarations or other traced payloads. /// Only the trace's copy is modified; transport messages retain their headers. fn redact_http_credentials(value: &mut serde_json::Value) { @@ -706,7 +733,45 @@ mod tests { use agent_client_protocol::RawJsonRpcMessage; use serde_json::json; - use super::{MessageInfo, Protocol, redact_http_credentials}; + use super::{MessageInfo, Protocol, ResponseEvent, redact_http_credentials, response_outcome}; + + #[test] + fn traced_mcp_outcomes_preserve_error_domain() { + let error = json!({"code":-32000,"message":"peer error","data":null,"extension":true}); + assert_eq!( + response_outcome(Protocol::Mcp, false, json!({"error":error})), + (Some(Protocol::Mcp), error.clone()) + ); + assert_eq!( + response_outcome(Protocol::Mcp, true, error.clone()), + (Some(Protocol::Acp), error) + ); + for result in [ + json!(null), + json!({"resultType":"input_required","requestState":"opaque"}), + ] { + assert_eq!( + response_outcome(Protocol::Mcp, false, json!({"result":result})), + (None, result) + ); + } + let acp_result = json!({"result": "not an MCP carrier"}); + assert_eq!( + response_outcome(Protocol::Acp, false, acp_result.clone()), + (None, acp_result) + ); + } + + #[test] + fn older_response_traces_without_error_domain_still_deserialize() { + let response: ResponseEvent = serde_json::from_value(json!({ + "ts": 0.0, "from": "Client", "to": "Agent", "id": 1, + "is_error": true, "payload": {"code": -32602, "message": "invalid"} + })) + .unwrap(); + assert!(response.is_error); + assert_eq!(response.error_domain, None); + } #[test] fn trace_credentials_are_redacted_in_nested_header_shapes() { diff --git a/src/agent-client-protocol-conductor/tests/mcp_cleanup_ownership.rs b/src/agent-client-protocol-conductor/tests/mcp_cleanup_ownership.rs new file mode 100644 index 00000000..a237c1bf --- /dev/null +++ b/src/agent-client-protocol-conductor/tests/mcp_cleanup_ownership.rs @@ -0,0 +1,294 @@ +#![cfg(feature = "unstable_protocol_v2")] + +//! Keep the tool runner unpolled during cancellation, while ACP still dispatches. +//! This proves that service completion alone cannot release request admission. + +use std::{ + future::Future, + pin::pin, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, + task::Poll, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, Client, ConnectionTo, Error, Responder, RunWithConnectionTo, V2ConnectionTo, + mcp_server::{ + McpConnectionTo, McpOutcome, McpRequest, McpRequestContext, McpServer, McpService, McpTool, + }, + schema::{ProtocolVersion, v2}, +}; +use futures::{FutureExt, future::BoxFuture, task::AtomicWaker}; +use schemars::JsonSchema; +use serde::{Deserialize, Serialize}; +use serde_json::json; +use tokio::sync::oneshot; + +#[derive(Deserialize, Serialize, JsonSchema)] +struct Input { + label: String, +} + +#[derive(Default)] +struct Gate { + paused: AtomicBool, + waker: AtomicWaker, +} + +impl Gate { + fn release(&self) { + self.paused.store(false, Ordering::Release); + self.waker.wake(); + } +} + +struct PausedRunner { + runner: R, + gate: Arc, +} + +impl> RunWithConnectionTo for PausedRunner { + async fn run_with_connection_to(self, cx: ConnectionTo) -> Result<(), Error> { + let mut running = pin!(self.runner.run_with_connection_to(cx)); + futures::future::poll_fn(|cx| { + self.gate.waker.register(cx.waker()); + if self.gate.paused.load(Ordering::Acquire) { + Poll::Pending + } else { + running.as_mut().poll(cx) + } + }) + .await + } +} + +struct ReleaseOnDrop(Arc); +impl Drop for ReleaseOnDrop { + fn drop(&mut self) { + self.0.release(); + } +} + +struct SignalOnDrop(Option>); +impl Drop for SignalOnDrop { + fn drop(&mut self) { + if let Some(tx) = self.0.take() { + let _sent = tx.send(()); + } + } +} + +struct ToolService { + tool: Arc, + finished: Arc>>>, +} + +impl McpService for ToolService +where + T: McpTool + 'static, +{ + fn execute( + &self, + request: McpRequest, + cx: McpRequestContext, + ) -> BoxFuture<'static, Result> { + let tool = self.tool.clone(); + let finished = self.finished.clone(); + Box::pin(async move { + let input: Input = + serde_json::from_value(request.params.expect("parameters")["arguments"].clone()) + .map_err(Error::into_internal_error)?; + let _finished = SignalOnDrop( + (input.label == "held") + .then(|| finished.lock().unwrap().take()) + .flatten(), + ); + let result = tokio::select! { + biased; + () = cx.operation_cancellation().cancelled() => Err(Error::request_cancelled()), + result = tool.call_tool(input, cx.connection().clone()) => result, + }?; + Ok(McpOutcome::Result(json!({ + "resultType": "complete", + "content": [{"type":"text", "text":result}] + }))) + }) + } +} + +async fn exercise( + tool: T, + runner: R, + started: oneshot::Receiver<()>, + dropped: oneshot::Receiver<()>, +) -> Result<(), Error> +where + T: McpTool + 'static, + R: RunWithConnectionTo + 'static, +{ + let gate = Arc::new(Gate::default()); + let (finished_tx, finished_rx) = oneshot::channel(); + let (result_tx, result_rx) = oneshot::channel(); + let state = Arc::new(Mutex::new(Some((started, dropped, finished_rx, result_tx)))); + let agent_gate = gate.clone(); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx| { + responder.respond(v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("cleanup-agent", "1"), + )) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected native declaration"); + }; + let server_id = server.server_id.clone(); + let (started, mut dropped, finished, result_tx) = + state.lock().unwrap().take().unwrap(); + let gate = agent_gate.clone(); + let work_cx = cx.clone(); + cx.spawn(async move { + let result = + async { + let _release_on_failure = ReleaseOnDrop(gate.clone()); + let request = + |label: &str| { + v2::MessageMcpRequest::new( + server_id.clone(), "same-logical-id", "tools/call", + ).params(json!({ + "name":"tool", "arguments":{"label":label}, + "_meta": { + "io.modelcontextprotocol/protocolVersion":"2026-07-28", + "io.modelcontextprotocol/clientCapabilities":{} + } + }).as_object().unwrap().clone()) + }; + let held = work_cx.send_request(request("held")); + started.await.map_err(Error::into_internal_error)?; + gate.paused.store(true, Ordering::Release); + held.cancel()?; + finished.await.map_err(Error::into_internal_error)?; + let mut response = Box::pin(held.block_task()); + + assert!( + matches!( + dropped.try_recv(), + Err(oneshot::error::TryRecvError::Empty) + ), + "runner remains paused" + ); + assert!( + response.as_mut().now_or_never().is_none(), + "cleanup precedes response" + ); + let duplicate = work_cx + .send_request(request("duplicate")) + .block_task() + .await + .expect_err("ID remains admitted during cleanup"); + assert_eq!(i32::from(duplicate.code), -32602); + + gate.release(); + let error = response.await.expect_err("cancelled operation"); + assert_eq!(i32::from(error.code), -32800); + assert!(dropped.try_recv().is_ok(), "tool dropped before reply"); + let healthy = + work_cx.send_request(request("after")).block_task().await?; + let v2::MessageMcpResponse::Result { result, .. } = healthy else { + panic!("ID reuse after cleanup should succeed"); + }; + assert_eq!(result["content"][0]["text"], "after"); + Ok::<_, Error>(()) + } + .await; + let _sent = result_tx.send(result); + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new("cleanup-session")) + }, + agent_client_protocol::on_receive_request!(), + ); + tokio::time::timeout( + Duration::from_secs(10), + Client.v2().connect_with(agent, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("cleanup-client", "1"), + )) + .block_task() + .await?; + let server = McpServer::new_service( + ToolService { + tool: Arc::new(tool), + finished: Arc::new(Mutex::new(Some(finished_tx))), + }, + "cleanup", + PausedRunner { runner, gate }, + ); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)? + .start_session() + .block_task() + .await?; + result_rx.await.map_err(Error::into_internal_error)? + }), + ) + .await + .expect("cleanup ownership regression timed out") +} + +#[tokio::test] +async fn mutable_tool_cleanup_precedes_id_release() -> Result<(), Error> { + let (started_tx, started_rx) = oneshot::channel(); + let (dropped_tx, dropped_rx) = oneshot::channel(); + let mut signals = Some((started_tx, dropped_tx)); + let (tool, runner) = agent_client_protocol::mcp_server::tool_fn_mut( + "tool", + "cleanup probe", + async move |input: Input, _cx: McpConnectionTo| { + if input.label == "held" { + let (started, dropped) = signals.take().unwrap(); + let _drop = SignalOnDrop(Some(dropped)); + let _sent = started.send(()); + std::future::pending::<()>().await; + } + Ok(input.label) + }, + agent_client_protocol::tool_fn_mut!(), + ); + exercise(tool, runner, started_rx, dropped_rx).await +} + +#[tokio::test] +async fn concurrent_tool_cleanup_precedes_id_release() -> Result<(), Error> { + let (started_tx, started_rx) = oneshot::channel(); + let (dropped_tx, dropped_rx) = oneshot::channel(); + let signals = Mutex::new(Some((started_tx, dropped_tx))); + let (tool, runner) = agent_client_protocol::mcp_server::tool_fn( + "tool", + "cleanup probe", + async move |input: Input, _cx: McpConnectionTo| { + if input.label == "held" { + let (started, dropped) = signals.lock().unwrap().take().unwrap(); + let _drop = SignalOnDrop(Some(dropped)); + let _sent = started.send(()); + std::future::pending::<()>().await; + } + Ok(input.label) + }, + agent_client_protocol::tool_fn!(), + ); + exercise(tool, runner, started_rx, dropped_rx).await +} diff --git a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs index 41927ada..e2d936d0 100644 --- a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs +++ b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs @@ -73,7 +73,7 @@ impl ConnectTo for NativeMcpProvider { assert_eq!(request.server_id.to_string(), SERVER_ID); self.request_count.fetch_add(1, Ordering::SeqCst); responder.respond(serde_json::from_value::( - serde_json::json!({"tools": []}), + serde_json::json!({"result":{"tools": []}}), )?) }, agent_client_protocol::on_receive_request!(), @@ -147,13 +147,18 @@ fn native_server() -> McpServer { } async fn http_post(url: &str, bearer: &str, id: i64) -> serde_json::Value { - let address = url.strip_prefix("http://").unwrap(); + let (address, route) = url + .strip_prefix("http://") + .unwrap() + .split_once('/') + .unwrap(); let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); let body = serde_json::json!({"jsonrpc":"2.0","id":id,"method":"tools/list", - "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28"}}}) + "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28", + "io.modelcontextprotocol/clientCapabilities":{}}}}) .to_string(); let request = format!( - "POST / HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: tools/list\r\nContent-Length: {}\r\n\r\n{body}", + "POST /{route} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: tools/list\r\nContent-Length: {}\r\n\r\n{body}", body.len() ); stream.write_all(request.as_bytes()).await.unwrap(); diff --git a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs index a7d10dc9..d28197a0 100644 --- a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs +++ b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs @@ -119,13 +119,23 @@ impl ConnectTo for TestProvider { }}} ]}), "tools/call" => serde_json::json!({"content":[]}), + "tools/error" => { + return responder.respond(serde_json::from_value::< + v2::MessageMcpResponse, + >( + serde_json::json!({"error":{"code":-32000,"message":"peer-owned", + "data":{"source":"backend"}}}), + )?); + } _ => { return responder.respond_with_error( agent_client_protocol::Error::method_not_found(), ); } }; - responder.respond(serde_json::from_value::(result)?) + responder.respond(serde_json::from_value::( + serde_json::json!({"result":result}), + )?) }, agent_client_protocol::on_receive_request!(), ) @@ -183,10 +193,15 @@ fn initialize() -> v2::InitializeRequest { } async fn post(url: &str, bearer: &str, method: &str, tool: &str) -> serde_json::Value { - let address = url.strip_prefix("http://").unwrap(); + let (address, route) = url + .strip_prefix("http://") + .unwrap() + .split_once('/') + .unwrap(); let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); let mut params = serde_json::json!({ - "_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28"} + "_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28", + "io.modelcontextprotocol/clientCapabilities":{}} }); if method == "tools/call" { params["name"] = serde_json::json!(tool); @@ -201,7 +216,7 @@ async fn post(url: &str, bearer: &str, method: &str, tool: &str) -> serde_json:: String::new() }; let request = format!( - "POST / HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: {method}\r\n{name}Content-Length: {}\r\n\r\n{body}", + "POST /{route} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: {method}\r\n{name}Content-Length: {}\r\n\r\n{body}", body.len() ); stream.write_all(request.as_bytes()).await.unwrap(); @@ -258,15 +273,20 @@ async fn modern_http_v2_requests_are_stateless_and_isolated() assert_eq!(server.headers[0].name, "Authorization"); (server.url.clone(), server.headers[0].value.clone()) }; - // No prior client tools/list: the adapter looks up the descriptor - // internally, rejecting annotated tools instead of skipping mirrors. + // A direct call does not require discovery or an internal tools/list lookup. let direct = post(&url, &bearer, "tools/call", "ping").await; assert_eq!( direct, serde_json::json!({"jsonrpc":"2.0","id":"same","result":{"content":[]}}) ); let annotated = post(&url, &bearer, "tools/call", "restricted").await; - assert_eq!(annotated["error"]["code"], -32602); + assert_eq!(annotated["result"], serde_json::json!({"content":[]})); + let peer_error = post(&url, &bearer, "tools/error", "").await; + assert_eq!( + peer_error["error"], + serde_json::json!({"code":-32000, + "message":"peer-owned","data":{"source":"backend"}}) + ); let (a, b) = tokio::join!( post(&url, &bearer, "tools/list", ""), post(&url, &bearer, "tools/list", "") @@ -274,7 +294,10 @@ async fn modern_http_v2_requests_are_stateless_and_isolated() assert_eq!( a, serde_json::json!({"jsonrpc":"2.0","id":"same","result":{"tools":[ - {"name":"ping","inputSchema":{"type":"object","properties":{}}} + {"name":"ping","inputSchema":{"type":"object","properties":{}}}, + {"name":"restricted","inputSchema":{"type":"object","properties":{ + "region":{"type":"string"} + }}} ]}}) ); assert_eq!(a, b); @@ -354,14 +377,15 @@ async fn closing_subscription_stream_cancels_only_its_native_request() let v2::McpServer::Http(server) = &observed[0] else { panic!("expected HTTP endpoint") }; (server.url.clone(), server.headers[0].value.clone()) }; - let address = url.strip_prefix("http://").unwrap(); + let (address, route) = url.strip_prefix("http://").unwrap().split_once('/').unwrap(); let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); let body = serde_json::json!({ "jsonrpc":"2.0","id":73,"method":"subscriptions/listen", - "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28"}} + "params":{"notifications":{"toolsListChanged":true}, "_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28", + "io.modelcontextprotocol/clientCapabilities":{}}} }).to_string(); let request = format!( - "POST / HTTP/1.1\r\nHost: {address}\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: subscriptions/listen\r\nContent-Length: {}\r\n\r\n{body}", + "POST /{route} HTTP/1.1\r\nHost: {address}\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: subscriptions/listen\r\nContent-Length: {}\r\n\r\n{body}", body.len() ); stream.write_all(request.as_bytes()).await.unwrap(); @@ -405,7 +429,7 @@ async fn closing_subscription_stream_cancels_only_its_native_request() } }).await.expect("closing the second stream must cancel its own ACP request"); let overflow = post(&url, &bearer, "subscriptions/flood", "").await; - assert_eq!(overflow["error"]["code"], -32000, "{overflow}"); + assert_eq!(overflow["error"]["code"], -33000, "{overflow}"); tokio::time::timeout(std::time::Duration::from_secs(3), async { while cancelled.load(Ordering::SeqCst) != 3 { tokio::task::yield_now().await; diff --git a/src/agent-client-protocol-conductor/tests/stateless_mcp_http.rs b/src/agent-client-protocol-conductor/tests/stateless_mcp_http.rs new file mode 100644 index 00000000..5532daf3 --- /dev/null +++ b/src/agent-client-protocol-conductor/tests/stateless_mcp_http.rs @@ -0,0 +1,311 @@ +//! Real rmcp client -> HTTP polyfill -> ACP/conductor -> rmcp server. + +use std::{ + future::Future, + path::PathBuf, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, Client, Error, + mcp_server::McpServer, + schema::{ + ProtocolVersion, + v1::{ + AgentCapabilities, InitializeRequest, InitializeResponse, McpCapabilities, + McpServer as AcpMcpServer, NewSessionRequest, NewSessionResponse, SessionCapabilities, + }, + }, +}; +use agent_client_protocol_conductor::{ConductorImpl, ProxiesAndAgent}; +use agent_client_protocol_polyfill::mcp_over_acp::McpOverAcpPolyfill; +use agent_client_protocol_rmcp::McpServerExt as _; +use rmcp::{ + ErrorData, RoleServer, ServerHandler, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, ClientCapabilities, ClientConfig, + Implementation, InputRequiredResult, ProtocolVersion as McpVersion, ServerCapabilities, + ServerConfig, SubscriptionFilter, Tool, ToolAnnotations, + }, + service::{ClientLifecycleMode, ClientServiceExt, RequestContext, SubscriptionContext}, + transport::{ + StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig, + }, +}; +use serde_json::{Value, json}; +use tokio::sync::mpsc; + +const TIMEOUT: Duration = Duration::from_secs(15); +const STATE: &str = "opaque/http/retry?keep=exact"; + +struct RealService { + listening: mpsc::UnboundedSender<()>, + stopped: mpsc::UnboundedSender<()>, + lists: Arc, +} + +struct NotifyStopped(mpsc::UnboundedSender<()>); + +impl Drop for NotifyStopped { + fn drop(&mut self) { + let _ = self.0.send(()); + } +} + +impl ServerHandler for RealService { + fn get_info(&self) -> ServerConfig { + ServerConfig::new( + ServerCapabilities::builder() + .enable_tools() + .enable_tool_list_changed() + .build(), + ) + } + + fn list_tools( + &self, + _request: Option, + _cx: RequestContext, + ) -> impl Future> + Send { + self.lists.fetch_add(1, Ordering::SeqCst); + let schema = json!({"type": "object"}).as_object().unwrap().clone(); + let annotated_schema = json!({ + "type": "object", + "properties": {"region": {"type": "string", "x-mcp-header": "Region"}} + }) + .as_object() + .unwrap() + .clone(); + std::future::ready(Ok(rmcp::model::ListToolsResult::with_all_items(vec![ + Tool::new("retry", "MRTR round trip", schema.clone()), + Tool::new("annotated", "Direct call", annotated_schema).with_annotations( + ToolAnnotations::from_raw(Some("Annotated".into()), Some(true), None, None, None), + ), + ]))) + } + + fn call_tool( + &self, + request: CallToolRequestParams, + cx: RequestContext, + ) -> impl Future> + Send { + std::future::ready( + match (request.name.as_ref(), request.request_state.as_deref()) { + ("retry", None) => { + let inputs = serde_json::from_value(json!({"confirmation": { + "method": "elicitation/create", + "params": {"mode": "form", "message": "Confirm", + "requestedSchema": {"type": "object", + "properties": {"approved": {"type": "boolean"}}}} + }})) + .expect("valid input request"); + Ok(InputRequiredResult::new(Some(inputs), Some(STATE.into())).into()) + } + ("retry", Some(STATE)) => Ok(CallToolResult::structured(json!({ + "state": request.request_state.clone(), + "responses": request.input_responses, + "marker": cx.meta.get("example/marker"), + })) + .into()), + ("annotated", None) => { + Ok(CallToolResult::structured(json!({"direct": true})).into()) + } + _ => Err(ErrorData::invalid_params( + "unknown tool or retry state", + None, + )), + }, + ) + } + + fn accepted_subscription_filter( + &self, + requested: &SubscriptionFilter, + ) -> Option { + Some(requested.clone()) + } + + async fn listen(&self, cx: SubscriptionContext) -> Result<(), ErrorData> { + // The owned adapter may cancel by dropping this future before it polls + // cx.cancelled() again. Observe actual cleanup, not a cooperative branch. + let _stopped = NotifyStopped(self.stopped.clone()); + cx.sink() + .notify_tool_list_changed() + .await + .map_err(|e| ErrorData::internal_error(e.to_string(), None))?; + let _ = self.listening.send(()); + cx.cancelled().await; + Ok(()) + } +} + +fn marked_call(name: &str, marker: &str) -> CallToolRequestParams { + let mut params = CallToolRequestParams::new(name.to_owned()); + params.meta = Some( + serde_json::from_value(json!({ + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {"elicitation": {"form": {}}}, + "io.modelcontextprotocol/clientInfo": {"name": "http-integration", "version": "1"}, + "example/marker": marker, + })) + .expect("valid request metadata"), + ); + params +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn real_rmcp_stateless_http_survives_subscription_cancellation() -> Result<(), Error> { + tokio::time::timeout(TIMEOUT, async { + let (endpoint_tx, mut endpoint_rx) = mpsc::unbounded_channel(); + let (listening_tx, mut listening_rx) = mpsc::unbounded_channel(); + let (stopped_tx, mut stopped_rx) = mpsc::unbounded_channel(); + let lists = Arc::new(AtomicUsize::new(0)); + let agent = Agent + .builder() + .on_receive_request( + async |request: InitializeRequest, responder, _cx| { + responder.respond( + InitializeResponse::new(request.protocol_version).agent_capabilities( + AgentCapabilities::new() + .session_capabilities(SessionCapabilities::new()) + .mcp_capabilities(McpCapabilities::new().http(true)), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: NewSessionRequest, responder, _cx| { + let [AcpMcpServer::Http(server)] = request.mcp_servers.as_slice() else { + panic!("expected a single HTTP MCP server declaration") + }; + assert_eq!(server.name, "real-rmcp"); + assert_eq!(server.headers.len(), 1); + assert_eq!(server.headers[0].name, "Authorization"); + endpoint_tx + .send((server.url.clone(), server.headers[0].value.clone())) + .expect("client still waiting for HTTP declaration"); + responder.respond(NewSessionResponse::new("real-http-session")) + }, + agent_client_protocol::on_receive_request!(), + ); + + Client.builder().connect_with( + ConductorImpl::new_agent( + "http-bridge", + ProxiesAndAgent::new(agent).proxy(McpOverAcpPolyfill::http()), + ), + async move |cx| { + cx.send_request(InitializeRequest::new(ProtocolVersion::V1)) + .block_task().await?; + let service = Arc::new(RealService { + listening: listening_tx, + stopped: stopped_tx, + lists: lists.clone(), + }); + cx.build_session(PathBuf::from("/tmp")) + .with_mcp_server(McpServer::::from_rmcp( + "real-rmcp", move || service.clone(), + ))? + .block_task() + .run_until(async move |_session| { + let (url, bearer) = endpoint_rx.recv().await.expect("HTTP declaration"); + let headers = [( + "Authorization".parse().expect("header name"), + bearer.parse().expect("header value"), + )].into_iter().collect(); + let transport = StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(url) + .custom_headers(headers), + ); + let config = ClientConfig::new( + serde_json::from_value::( + json!({"elicitation": {"form": {}}}), + ).expect("valid capabilities"), + Implementation::new("http-integration", "1"), + ).with_protocol_version(McpVersion::V_2026_07_28); + let client = config.serve_with_lifecycle( + transport, + ClientLifecycleMode::Discover { + preferred_versions: vec![McpVersion::V_2026_07_28], + }, + ).await.map_err(Error::into_internal_error)?; + + // A direct call before any list also tests absence of hidden lists. + let first = client.call_tool_once(marked_call("retry", "first")) + .await.map_err(Error::into_internal_error)?; + let CallToolResponse::InputRequired(first) = first else { + panic!("expected input_required, got {first:?}"); + }; + assert_eq!(lists.load(Ordering::SeqCst), 0, "no hidden tools/list"); + assert_eq!(first.request_state.as_deref(), Some(STATE)); + assert_eq!( + serde_json::to_value(&first.input_requests) + .map_err(Error::into_internal_error)?["confirmation"]["method"], + "elicitation/create" + ); + let responses: Value = + json!({"confirmation": {"action": "accept", "content": {"approved": true}}}); + let second = client.call_tool_once( + marked_call("retry", "second") + .with_request_state(first.request_state.expect("opaque state")) + .with_input_responses(serde_json::from_value(responses.clone()) + .map_err(Error::into_internal_error)?), + ).await.map_err(Error::into_internal_error)?; + let CallToolResponse::Complete(second) = second else { + panic!("expected completed retry, got {second:?}"); + }; + assert_eq!(second.structured_content.as_ref().unwrap()["state"], STATE); + assert_eq!(second.structured_content.as_ref().unwrap()["responses"], responses); + assert_eq!(second.structured_content.as_ref().unwrap()["marker"], "second"); + + let filter = SubscriptionFilter::builder().tools_list_changed().build(); + let mut subscription = client.listen(filter.clone()).await + .map_err(Error::into_internal_error)?; + listening_rx.recv().await.expect("subscription service started"); + assert_eq!(subscription.acknowledged(), &filter); + let notification = subscription.next().await + .map_err(Error::into_internal_error)? + .expect("filtered notification"); + let notification_json = serde_json::to_value(¬ification) + .map_err(Error::into_internal_error)?; + assert_eq!(notification_json["method"], "notifications/tools/list_changed"); + assert_eq!( + notification_json["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"], + serde_json::to_value(subscription.id()) + .map_err(Error::into_internal_error)? + ); + + // An active HTTP SSE listen must not block an ordinary POST. + let parallel = client.call_tool_once(marked_call("annotated", "parallel")) + .await.map_err(Error::into_internal_error)?; + assert!(matches!(parallel, CallToolResponse::Complete(_))); + subscription.cancel().await.map_err(Error::into_internal_error)?; + stopped_rx.recv().await.expect("subscription cancelled upstream"); + drop(subscription); + let direct = client.call_tool_once(marked_call("annotated", "after-cancel")) + .await.map_err(Error::into_internal_error)?; + assert!(matches!(direct, CallToolResponse::Complete(_))); + let tools = client.list_tools(None).await.map_err(Error::into_internal_error)?; + let annotated = tools.tools.iter().find(|tool| tool.name == "annotated") + .expect("annotated tool still listed"); + assert_eq!(annotated.annotations.as_ref().unwrap().read_only_hint, Some(true)); + assert_eq!( + annotated.input_schema.get("properties").unwrap()["region"], + json!({"type": "string"}) + ); + assert_eq!(lists.load(Ordering::SeqCst), 1); + client.cancel().await.map_err(Error::into_internal_error)?; + Ok(()) + }) + .await + }, + ).await + }) + .await + .expect("rmcp/HTTP/ACP integration timed out") +} diff --git a/src/agent-client-protocol-conductor/tests/test_tool_fn.rs b/src/agent-client-protocol-conductor/tests/test_tool_fn.rs index 67f91609..0d601fb9 100644 --- a/src/agent-client-protocol-conductor/tests/test_tool_fn.rs +++ b/src/agent-client-protocol-conductor/tests/test_tool_fn.rs @@ -79,3 +79,145 @@ async fn test_tool_fn_greet() -> Result<(), agent_client_protocol::Error> { Ok(()) } + +/// A cancelled call must not poison the mutable runner, and queued work whose +/// result receiver has gone away must never enter the user's closure. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn cancelled_tool_fn_mut_keeps_acp_alive() -> Result<(), agent_client_protocol::Error> { + use agent_client_protocol::{ + Agent, Client, Error, Responder, V2ConnectionTo, + schema::{ProtocolVersion, v2}, + }; + use std::{ + sync::{Arc, Mutex}, + time::Duration, + }; + use tokio::sync::oneshot; + + #[derive(Debug, Deserialize, Serialize, JsonSchema)] + struct Input { + name: String, + } + let (started_tx, started_rx) = oneshot::channel(); + let started = Arc::new(Mutex::new(Some(started_tx))); + let calls = Arc::new(Mutex::new(Vec::::new())); + let (result_tx, result_rx) = oneshot::channel(); + let invocation = Arc::new(Mutex::new(Some((started_rx, result_tx)))); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("runner-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new() + .session(v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + )), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let [v2::McpServer::Acp(declaration)] = request.mcp_servers.as_slice() else { + panic!("expected one ACP MCP server") + }; + let server = declaration.server_id.clone(); + let (started_rx, result_tx) = + invocation.lock().unwrap().take().expect("one session"); + let call_cx = cx.clone(); + cx.spawn(async move { + let result = async { + let make_request = |id: &str, name: &str| { + let params = serde_json::json!({ + "name": "hold", + "arguments": {"name": name}, + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } + }); + v2::MessageMcpRequest::new(server.clone(), id.to_owned(), "tools/call") + .params(params.as_object().expect("object params").clone()) + }; + let running = call_cx.send_request(make_request("running", "running")); + started_rx.await.map_err(Error::into_internal_error)?; + let queued = call_cx.send_request(make_request("queued", "queued")); + tokio::task::yield_now().await; + queued.cancel()?; + running.cancel()?; + for request in [running, queued] { + let failure = + request.block_task().await.expect_err("cancelled request"); + assert_eq!(i32::from(failure.code), -32800); + } + let healthy = call_cx + .send_request(make_request("after", "after")) + .block_task() + .await?; + let v2::MessageMcpResponse::Result { result, .. } = healthy else { + panic!("healthy tool call did not produce an MCP result") + }; + assert_eq!(result["isError"], false, "healthy result: {result}"); + Ok::<_, Error>(()) + } + .await; + let _sent = result_tx.send(result); + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new( + "runner-session", + ))) + }, + agent_client_protocol::on_receive_request!(), + ); + tokio::time::timeout( + Duration::from_secs(10), + Client.v2().connect_with(agent, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("runner-client", "1"), + )) + .block_task() + .await?; + let recorded = calls.clone(); + let started = started.clone(); + let server = McpServer::::builder("runner") + .tool_fn_mut( + "hold", + "Hold a mutable runner", + async move |input: Input, _cx| { + recorded.lock().unwrap().push(input.name.clone()); + if input.name == "running" { + if let Some(tx) = started.lock().unwrap().take() { + let _sent = tx.send(()); + } + std::future::pending::<()>().await; + } + Ok(serde_json::json!({"value": input.name})) + }, + agent_client_protocol::tool_fn_mut!(), + ) + .build(); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)? + .start_session() + .block_task() + .await?; + result_rx.await.map_err(Error::into_internal_error)??; + assert_eq!(*calls.lock().unwrap(), ["running", "after"]); + Ok::<_, Error>(()) + }), + ) + .await + .expect("cancelled MCP tool call timed out") +} diff --git a/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs b/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs index 11c51fed..1582ba90 100644 --- a/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs +++ b/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs @@ -367,6 +367,7 @@ async fn test_trace_client_mcp_server() -> Result<(), agent_client_protocol::Err to: "Client", id: String("id:0"), is_error: false, + error_domain: None, payload: Object { "protocolVersion": Number(1), "agentCapabilities": Object { @@ -430,6 +431,7 @@ async fn test_trace_client_mcp_server() -> Result<(), agent_client_protocol::Err to: "Client", id: String("id:1"), is_error: false, + error_domain: None, payload: Object { "sessionId": String("session:0"), "modes": Object { @@ -521,6 +523,7 @@ async fn test_trace_client_mcp_server() -> Result<(), agent_client_protocol::Err to: "Client", id: String("id:2"), is_error: false, + error_domain: None, payload: Object { "stopReason": String("end_turn"), }, diff --git a/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs b/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs index f8318573..9bede599 100644 --- a/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs +++ b/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs @@ -384,6 +384,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Proxy(0)", id: String("id:1"), is_error: false, + error_domain: None, payload: Object { "protocolVersion": Number(1), "agentCapabilities": Object { @@ -426,6 +427,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Client", id: String("id:0"), is_error: false, + error_domain: None, payload: Object { "protocolVersion": Number(1), "agentCapabilities": Object { @@ -504,6 +506,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Proxy(0)", id: String("id:3"), is_error: false, + error_domain: None, payload: Object { "sessionId": String("session:0"), "modes": Object { @@ -554,6 +557,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Client", id: String("id:2"), is_error: false, + error_domain: None, payload: Object { "sessionId": String("session:0"), "modes": Object { @@ -665,6 +669,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Proxy(1)", id: String("id:6"), is_error: false, + error_domain: None, payload: Object { "resultType": String("complete"), "supportedVersions": Array [ @@ -692,78 +697,6 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> from: "Proxy(1)", to: "Proxy(0)", id: String("id:7"), - method: "tools/list", - session: None, - params: Object { - "_meta": Object { - "io.modelcontextprotocol/protocolVersion": String("2026-07-28"), - "io.modelcontextprotocol/clientInfo": Object { - "name": String("testy"), - "version": String("0.11.0"), - }, - "io.modelcontextprotocol/clientCapabilities": Object {}, - "progressToken": Number(0), - }, - }, - }, - ), - Response( - ResponseEvent { - ts: 0.0, - from: "Proxy(0)", - to: "Proxy(1)", - id: String("id:7"), - is_error: false, - payload: Object { - "resultType": String("complete"), - "ttlMs": Number(0), - "cacheScope": String("private"), - "tools": Array [ - Object { - "name": String("echo"), - "description": String("Echoes back the input message"), - "inputSchema": Object { - "$schema": String("https://json-schema.org/draft/2020-12/schema"), - "title": String("EchoParams"), - "description": String("Parameters for the echo tool"), - "type": String("object"), - "properties": Object { - "message": Object { - "description": String("The message to echo back"), - "type": String("string"), - }, - }, - "required": Array [ - String("message"), - ], - }, - "outputSchema": Object { - "$schema": String("https://json-schema.org/draft/2020-12/schema"), - "title": String("EchoOutput"), - "description": String("Output from the echo tool"), - "type": String("object"), - "properties": Object { - "result": Object { - "description": String("The echoed message"), - "type": String("string"), - }, - }, - "required": Array [ - String("result"), - ], - }, - }, - ], - }, - }, - ), - Request( - RequestEvent { - ts: 0.0, - protocol: Mcp, - from: "Proxy(1)", - to: "Proxy(0)", - id: String("id:8"), method: "tools/call", session: None, params: Object { @@ -788,8 +721,9 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> ts: 0.0, from: "Proxy(0)", to: "Proxy(1)", - id: String("id:8"), + id: String("id:7"), is_error: false, + error_domain: None, payload: Object { "resultType": String("complete"), "content": Array [ @@ -833,6 +767,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Proxy(0)", id: String("id:5"), is_error: false, + error_domain: None, payload: Object { "stopReason": String("end_turn"), }, @@ -866,6 +801,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Client", id: String("id:4"), is_error: false, + error_domain: None, payload: Object { "stopReason": String("end_turn"), }, diff --git a/src/agent-client-protocol-conductor/tests/trace_snapshot.rs b/src/agent-client-protocol-conductor/tests/trace_snapshot.rs index 51b48b51..d215ca1b 100644 --- a/src/agent-client-protocol-conductor/tests/trace_snapshot.rs +++ b/src/agent-client-protocol-conductor/tests/trace_snapshot.rs @@ -216,6 +216,7 @@ async fn test_trace_snapshot() -> Result<(), agent_client_protocol::Error> { to: "Client", id: String("id:0"), is_error: false, + error_domain: None, payload: Object { "protocolVersion": Number(1), "agentCapabilities": Object { @@ -273,6 +274,7 @@ async fn test_trace_snapshot() -> Result<(), agent_client_protocol::Error> { to: "Client", id: String("id:1"), is_error: false, + error_domain: None, payload: Object { "sessionId": String("session:0"), "modes": Object { @@ -364,6 +366,7 @@ async fn test_trace_snapshot() -> Result<(), agent_client_protocol::Error> { to: "Client", id: String("id:2"), is_error: false, + error_domain: None, payload: Object { "stopReason": String("end_turn"), }, diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index 1c899872..14b92f11 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -4,14 +4,15 @@ use std::{ }; use agent_client_protocol::{ - Agent, Channel, Client, ConnectTo, Error as AcpError, RawJsonRpcMessage, TransportBatchEntry, + Agent, BudgetedFrame, Channel, Client, ConnectTo, Error as AcpError, FrameAdmission, + FramePermit, FrameReceiver, FrameSender, RawJsonRpcMessage, TransportBatchEntry, TransportFrame, schema::v1::{RequestId, Response as RpcResponse}, }; use async_tungstenite::tungstenite::Message as WsMessage; use futures::{ - Stream, StreamExt, - channel::mpsc::{self, UnboundedSender}, + SinkExt, Stream, StreamExt, + channel::mpsc, future::{BoxFuture, FutureExt}, pin_mut, stream::FuturesUnordered, @@ -20,8 +21,9 @@ use thiserror::Error; use tracing::{debug, error, trace, warn}; use crate::protocol::{ - HEADER_CONNECTION_ID, HEADER_SESSION_ID, is_initialize_request, is_response_only_shape, - method_for_message, method_requires_session_header, session_id_from_message, + HEADER_CONNECTION_ID, HEADER_SESSION_ID, cancelled_request_id, is_initialize_request, + is_response_only_shape, method_for_message, method_requires_session_header, + session_id_from_message, }; #[derive(Debug, Error)] @@ -123,9 +125,12 @@ impl ConnectTo for HttpClient { } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), AcpError>>) { + fn into_channel_and_future(self) -> (Channel, agent_client_protocol::ConnectionDriver) { let (caller, transport) = Channel::duplex(); - (caller, Box::pin(run(self, transport))) + ( + caller, + agent_client_protocol::ConnectionDriver::new(run(self, transport)), + ) } } @@ -138,15 +143,18 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { rx: mut outgoing, tx: incoming, } = channel; - let (sse_event_tx, mut sse_event_rx) = mpsc::unbounded::(); + let admission = incoming.admission(); + let max_operations = admission.limits().max_queued_frames.max(1); + let (sse_event_tx, mut sse_event_rx) = mpsc::channel::(max_operations); let connection = HttpConnection::new(endpoint, http); let mut state = ClientState { connection: connection.clone(), open_session_streams: HashSet::new(), pending_requests: HashMap::new(), + pending_request_leases: HashMap::new(), incoming, }; - let mut lifecycle = HttpTransportLifecycle::new(connection); + let mut lifecycle = HttpTransportLifecycle::new(connection, admission, max_operations); let mut posts = PostQueues::default(); let mut buffered_outgoing = VecDeque::new(); let mut outgoing_closed = false; @@ -200,8 +208,8 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { let Some(event) = event else { continue; }; - let open_session_ids = state.sessions_to_open_for_responses(&event.frame); - state.deliver_frame(event.frame); + let open_session_ids = state.sessions_to_open_for_responses(event.frame.frame()); + state.deliver_budgeted(event.frame).await?; for session_id in open_session_ids { match lifecycle .start_sse( @@ -242,7 +250,9 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { } }; - let is_response_only = is_response_only_frame(&frame); + let bypass_ordered = + is_response_only_frame(frame.frame()) || is_cancellation_frame(frame.frame()); + let (frame, permit) = frame.into_parts(); let msg = match frame { TransportFrame::Single(message) => message, frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => { @@ -254,6 +264,7 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { // Response-only batches answer SSE-delivered callbacks and // must not be blocked behind the request they answer. Ok((post, session_ids)) => { + state.attach_pending_permits(&post.pending_requests, &permit); for session_id in session_ids { match lifecycle .start_sse( @@ -276,10 +287,15 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { Err(error) => break 'transport Err(error), } } - if is_response_only { - posts.responses.push(post); + if let Err(error) = + check_post_capacity(&posts, max_operations, bypass_ordered) + { + break 'transport Err(error); + } + if bypass_ordered { + posts.responses.push_budgeted(post, permit); } else { - posts.ordered.push(post); + posts.ordered.push_budgeted(post, permit); } } Err(error) => { @@ -356,11 +372,20 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { } } + if let Err(error) = check_post_capacity(&posts, max_operations, bypass_ordered) { + break Err(error); + } match state.prepare_post(msg) { - // Responses answer SSE-delivered callbacks and must not be blocked - // behind a POST that may be waiting for that callback response. - Ok(post) if is_response_only => posts.responses.push(post), - Ok(post) => posts.ordered.push(post), + // Responses and cancellation must not be blocked behind a POST + // that may itself be waiting for their delivery. + Ok(post) => { + state.attach_pending_permits(&post.pending_requests, &permit); + if bypass_ordered { + posts.responses.push_budgeted(post, permit); + } else { + posts.ordered.push_budgeted(post, permit); + } + } Err(e) => { error!("POST failed: {e}"); break Err(AcpError::internal_error().data(format!("POST: {e}"))); @@ -383,12 +408,32 @@ fn sse_setup_blocked_output_error() -> AcpError { .data("outgoing channel closed while accepted messages awaited SSE stream establishment") } +fn post_capacity_error() -> AcpError { + AcpError::internal_error().data("HTTP POST operation capacity exceeded") +} + +fn check_post_capacity( + posts: &PostQueues, + max_operations: usize, + bypass_ordered: bool, +) -> Result<(), AcpError> { + // Keep one operation available for callbacks/cancellation even while the + // ordered data POST is waiting for exactly such a response. + let reserved = usize::from(!bypass_ordered && max_operations > 1); + if posts.len() >= max_operations.saturating_sub(reserved).max(1) { + Err(post_capacity_error()) + } else { + Ok(()) + } +} + fn handle_completed_post( state: &mut ClientState, completed: CompletedPost, ) -> Result<(), AcpError> { let CompletedPost { pending_requests, + cancelled_requests, result, } = completed; if let Err(error) = result { @@ -396,6 +441,9 @@ fn handle_completed_post( error!("POST failed: {error}"); Err(AcpError::internal_error().data(format!("POST: {error}"))) } else { + for id in cancelled_requests { + state.cancel_pending_request(&id); + } Ok(()) } } @@ -403,8 +451,14 @@ fn handle_completed_post( fn queue_response_post( state: &mut ClientState, posts: &mut PostQueues, - frame: TransportFrame, + frame: BudgetedFrame, ) -> Result<(), AcpError> { + check_post_capacity( + posts, + state.incoming.admission().limits().max_queued_frames.max(1), + true, + )?; + let (frame, permit) = frame.into_parts(); let post = match frame { TransportFrame::Single(message) => state.prepare_post(message), frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => { @@ -418,7 +472,8 @@ fn queue_response_post( error!("POST failed: {error}"); AcpError::internal_error().data(format!("POST: {error}")) })?; - posts.responses.push(post); + state.attach_pending_permits(&post.pending_requests, &permit); + posts.responses.push_budgeted(post, permit); Ok(()) } @@ -441,8 +496,16 @@ fn is_response_only_frame(frame: &TransportFrame) -> bool { } } +fn is_cancellation_frame(frame: &TransportFrame) -> bool { + matches!( + frame, + TransportFrame::Single(RawJsonRpcMessage::Notification(message)) + if message.method.as_ref() == "$/cancel_request" + ) +} + enum HttpLoopEvent { - Outgoing(Option), + Outgoing(Option), SseEvent(Option), SseFailure(SseFailure), Post(CompletedPost), @@ -456,7 +519,7 @@ struct SseFailure { #[derive(Debug)] struct SseMessage { - frame: TransportFrame, + frame: BudgetedFrame, } #[derive(Clone, Debug)] @@ -535,6 +598,9 @@ impl HttpConnection { if let Err(e) = http .delete(endpoint) .header(HEADER_CONNECTION_ID, connection_id) + // A stalled peer must not keep the transport's shutdown (and its + // retained POST/SSE permits) alive indefinitely. + .timeout(std::time::Duration::from_secs(2)) .send() .await { @@ -547,6 +613,8 @@ impl HttpConnection { struct HttpTransportLifecycle { connection: HttpConnection, sse_tasks: SseTasks, + admission: FrameAdmission, + max_tasks: usize, } #[derive(Clone, Copy, Debug, Eq, PartialEq)] @@ -556,25 +624,27 @@ enum SseStartOutcome { } struct SseStartContext<'a> { - events: &'a mut mpsc::UnboundedReceiver, - outgoing: &'a mut mpsc::UnboundedReceiver, - buffered_outgoing: &'a mut VecDeque, + events: &'a mut mpsc::Receiver, + outgoing: &'a mut FrameReceiver, + buffered_outgoing: &'a mut VecDeque, posts: &'a mut PostQueues, state: &'a mut ClientState, } impl HttpTransportLifecycle { - fn new(connection: HttpConnection) -> Self { + fn new(connection: HttpConnection, admission: FrameAdmission, max_tasks: usize) -> Self { Self { connection, sse_tasks: SseTasks::default(), + admission, + max_tasks, } } async fn start_sse( &mut self, session_id: Option, - event_tx: UnboundedSender, + event_tx: mpsc::Sender, context: SseStartContext<'_>, ) -> Result { let SseStartContext { @@ -585,7 +655,7 @@ impl HttpTransportLifecycle { state, } = context; let mut establishing = FuturesUnordered::new(); - establishing.push(self.begin_sse(session_id, event_tx.clone())); + establishing.push(self.begin_sse(session_id, event_tx.clone())?); loop { if establishing.is_empty() { @@ -625,20 +695,29 @@ impl HttpTransportLifecycle { } SseStartWait::Failure(failure) => return Err(sse_failure_error(failure)), SseStartWait::SseEvent(Some(event)) => { - let open_session_ids = state.sessions_to_open_for_responses(&event.frame); - state.deliver_frame(event.frame); + let open_session_ids = + state.sessions_to_open_for_responses(event.frame.frame()); + state.deliver_budgeted(event.frame).await?; for session_id in open_session_ids { - establishing.push(self.begin_sse(Some(session_id), event_tx.clone())); + establishing.push(self.begin_sse(Some(session_id), event_tx.clone())?); } } SseStartWait::SseEvent(None) => { return Err(AcpError::internal_error().data("SSE event channel closed")); } SseStartWait::Post(completed) => handle_completed_post(state, completed)?, - SseStartWait::Outgoing(Some(frame)) if is_response_only_frame(&frame) => { + SseStartWait::Outgoing(Some(frame)) + if is_response_only_frame(frame.frame()) + || is_cancellation_frame(frame.frame()) => + { queue_response_post(state, posts, frame)?; } - SseStartWait::Outgoing(Some(frame)) => buffered_outgoing.push_back(frame), + SseStartWait::Outgoing(Some(frame)) => { + if buffered_outgoing.len() + posts.len() >= self.max_tasks { + return Err(post_capacity_error()); + } + buffered_outgoing.push_back(frame); + } SseStartWait::Outgoing(None) => return Ok(SseStartOutcome::OutgoingClosed), } } @@ -647,16 +726,20 @@ impl HttpTransportLifecycle { fn begin_sse( &mut self, session_id: Option, - event_tx: UnboundedSender, - ) -> futures::channel::oneshot::Receiver<()> { + event_tx: mpsc::Sender, + ) -> Result, AcpError> { + if self.sse_tasks.len() >= self.max_tasks { + return Err(AcpError::internal_error().data("HTTP SSE stream capacity exceeded")); + } let (established_tx, established_rx) = futures::channel::oneshot::channel(); self.sse_tasks.push(run_sse( self.connection.clone(), session_id, event_tx, established_tx, + self.admission.clone(), )); - established_rx + Ok(established_rx) } async fn next_sse_failure(&mut self) -> SseFailure { @@ -674,7 +757,7 @@ enum SseStartWait { Failure(SseFailure), SseEvent(Option), Post(CompletedPost), - Outgoing(Option), + Outgoing(Option), } impl Drop for HttpTransportLifecycle { @@ -687,15 +770,17 @@ impl Drop for HttpTransportLifecycle { fn run_sse( connection: HttpConnection, session_id: Option, - event_tx: UnboundedSender, + event_tx: mpsc::Sender, established_tx: futures::channel::oneshot::Sender<()>, + admission: FrameAdmission, ) -> BoxFuture<'static, SseFailure> { Box::pin(async move { let label = session_id.clone(); - let error = match read_sse(connection, session_id, event_tx, established_tx).await { - Ok(()) => "SSE stream closed".to_string(), - Err(e) => e, - }; + let error = + match read_sse(connection, session_id, event_tx, established_tx, admission).await { + Ok(()) => "SSE stream closed".to_string(), + Err(e) => e, + }; warn!(session_id = ?label, "SSE stream ended: {error}"); SseFailure { session_id: label, @@ -710,6 +795,10 @@ struct SseTasks { } impl SseTasks { + fn len(&self) -> usize { + self.handles.len() + } + fn push(&mut self, task: BoxFuture<'static, SseFailure>) { self.handles.push(task); } @@ -732,25 +821,31 @@ struct ClientState { connection: HttpConnection, open_session_streams: HashSet, pending_requests: HashMap>, - incoming: futures::channel::mpsc::UnboundedSender, + pending_request_leases: HashMap>, + incoming: FrameSender, } struct PendingPost { pending_requests: Vec<(RequestId, String)>, + cancelled_requests: Vec, response: BoxFuture<'static, Result<(), String>>, } impl PendingPost { - fn into_completion(self) -> BoxFuture<'static, CompletedPost> { + fn into_completion(self, permit: Option) -> BoxFuture<'static, CompletedPost> { let Self { pending_requests, + cancelled_requests, response, } = self; async move { - CompletedPost { + let completed = CompletedPost { pending_requests, + cancelled_requests, result: response.await, - } + }; + drop(permit); + completed } .boxed() } @@ -759,12 +854,13 @@ impl PendingPost { #[derive(Debug)] struct CompletedPost { pending_requests: Vec<(RequestId, String)>, + cancelled_requests: Vec, result: Result<(), String>, } #[derive(Default)] struct PostQueue { - queued: VecDeque, + queued: VecDeque<(PendingPost, Option)>, in_flight: Option>, } @@ -778,11 +874,25 @@ impl PostQueues { fn is_empty(&self) -> bool { self.ordered.is_empty() && self.responses.is_empty() } + + fn len(&self) -> usize { + self.ordered.len() + self.responses.len() + } } impl PostQueue { + fn len(&self) -> usize { + self.queued.len() + usize::from(self.in_flight.is_some()) + } + + #[cfg(test)] fn push(&mut self, post: PendingPost) { - self.queued.push_back(post); + self.queued.push_back((post, None)); + self.start_next(); + } + + fn push_budgeted(&mut self, post: PendingPost, permit: FramePermit) { + self.queued.push_back((post, Some(permit))); self.start_next(); } @@ -800,9 +910,9 @@ impl PostQueue { fn start_next(&mut self) { if self.in_flight.is_none() - && let Some(post) = self.queued.pop_front() + && let Some((post, permit)) = self.queued.pop_front() { - self.in_flight = Some(post.into_completion()); + self.in_flight = Some(post.into_completion(permit)); } } @@ -859,14 +969,14 @@ impl ClientState { message, RawJsonRpcMessage::Response(RpcResponse::Error { .. }) ) { - self.deliver(message); + self.deliver(message).await.map_err(|e| e.to_string())?; self.connection.close().await; return Ok(InitializeOutcome::Rejected); } connection_id .ok_or_else(|| format!("server did not return {HEADER_CONNECTION_ID} header"))?; - self.deliver(message); + self.deliver(message).await.map_err(|e| e.to_string())?; Ok(InitializeOutcome::Connected) } @@ -889,6 +999,8 @@ impl ClientState { let pending_requests = pending_request_for_message(&msg) .into_iter() .collect::>(); + let cancelled_requests = cancelled_request_id(&msg).into_iter().collect(); + self.check_pending_request_capacity(pending_requests.len())?; self.track_pending_requests(&pending_requests); let response = async move { @@ -902,6 +1014,7 @@ impl ClientState { }; Ok(PendingPost { pending_requests, + cancelled_requests, response: response.boxed(), }) } @@ -911,6 +1024,7 @@ impl ClientState { frame: TransportFrame, ) -> Result<(PendingPost, Vec), String> { let bookkeeping = FrameBookkeeping::for_frame(&frame)?; + self.check_pending_request_capacity(bookkeeping.pending_requests.len())?; let connection_id = self .connection .connection_id() @@ -937,6 +1051,7 @@ impl ClientState { Ok(( PendingPost { pending_requests: bookkeeping.pending_requests, + cancelled_requests: bookkeeping.cancelled_requests, response: response.boxed(), }, session_ids, @@ -952,16 +1067,43 @@ impl ClientState { } } + fn check_pending_request_capacity(&self, additional: usize) -> Result<(), String> { + let limit = self.incoming.admission().limits().max_queued_frames.max(1); + let existing: usize = self.pending_requests.values().map(VecDeque::len).sum(); + if additional > limit.saturating_sub(existing) { + Err("HTTP pending request capacity exceeded".to_string()) + } else { + Ok(()) + } + } + + fn attach_pending_permits( + &mut self, + pending_requests: &[(RequestId, String)], + permit: &FramePermit, + ) { + for (id, _) in pending_requests { + self.pending_request_leases + .entry(id.clone()) + .or_default() + .push_back(permit.clone()); + } + } + fn remove_pending_requests(&mut self, pending_requests: &[(RequestId, String)]) { for (id, method) in pending_requests.iter().rev() { let remove_entry = self.pending_requests.get_mut(id).is_some_and(|methods| { if let Some(index) = methods.iter().rposition(|candidate| candidate == method) { methods.remove(index); + if let Some(leases) = self.pending_request_leases.get_mut(id) { + leases.remove(index); + } } methods.is_empty() }); if remove_entry { self.pending_requests.remove(id); + self.pending_request_leases.remove(id); } } } @@ -971,12 +1113,30 @@ impl ClientState { let methods = self.pending_requests.get_mut(id)?; (methods.pop_front(), methods.is_empty()) }; + if let Some(leases) = self.pending_request_leases.get_mut(id) { + leases.pop_front(); + } if remove_entry { self.pending_requests.remove(id); + self.pending_request_leases.remove(id); } method } + fn cancel_pending_request(&mut self, id: &RequestId) { + let Some(methods) = self.pending_requests.get_mut(id) else { + return; + }; + methods.pop_front(); + if let Some(leases) = self.pending_request_leases.get_mut(id) { + leases.pop_front(); + } + if methods.is_empty() { + self.pending_requests.remove(id); + self.pending_request_leases.remove(id); + } + } + fn register_session_streams( &mut self, session_ids: impl IntoIterator, @@ -1031,14 +1191,20 @@ impl ClientState { } } - fn deliver(&self, msg: RawJsonRpcMessage) { - self.deliver_frame(TransportFrame::Single(msg)); + async fn deliver(&self, msg: RawJsonRpcMessage) -> Result<(), AcpError> { + self.deliver_frame(TransportFrame::Single(msg)).await } - fn deliver_frame(&self, frame: TransportFrame) { - if self.incoming.unbounded_send(frame).is_err() { - debug!("upstream channel closed; dropping inbound message"); - } + async fn deliver_frame(&self, frame: TransportFrame) -> Result<(), AcpError> { + self.incoming.send_frame(frame).await + } + + async fn deliver_budgeted(&self, frame: BudgetedFrame) -> Result<(), AcpError> { + self.incoming + .clone() + .send(frame) + .await + .map_err(|error| AcpError::internal_error().data(format!("deliver SSE frame: {error}"))) } } @@ -1046,6 +1212,7 @@ impl ClientState { struct FrameBookkeeping { session_ids: Vec, pending_requests: Vec<(RequestId, String)>, + cancelled_requests: Vec, } impl FrameBookkeeping { @@ -1074,6 +1241,8 @@ impl FrameBookkeeping { if let Some(pending_request) = pending_request_for_message(message) { self.pending_requests.push(pending_request); } + self.cancelled_requests + .extend(cancelled_request_id(message)); Ok(()) } } @@ -1096,8 +1265,9 @@ fn is_session_opening_method(method: &str) -> bool { async fn read_sse( connection: HttpConnection, session_id: Option, - event_tx: UnboundedSender, + mut event_tx: mpsc::Sender, established_tx: futures::channel::oneshot::Sender<()>, + admission: FrameAdmission, ) -> Result<(), String> { let connection_id = connection .connection_id() @@ -1117,16 +1287,45 @@ async fn read_sse( trace!(session_id = ?session_id, "SSE stream open"); let _ = established_tx.send(()); - let mut events = eventsource_stream::EventStream::new(response.bytes_stream()); + // Cap each event before EventStream buffers its data fields or JSON parsing + // materializes the payload. A blank line terminates one SSE event. + let max_frame_bytes = admission.limits().max_frame_bytes; + let mut event_bytes = 0usize; + let mut line_has_data = false; + let mut events = + eventsource_stream::EventStream::new(response.bytes_stream().map(move |chunk| { + let chunk = chunk.map_err(std::io::Error::other)?; + for &byte in &chunk { + event_bytes += 1; + if event_bytes > max_frame_bytes { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "SSE event exceeds maximum JSON-RPC frame size", + )); + } + if byte == b'\n' { + if !line_has_data { + event_bytes = 0; + } + line_has_data = false; + } else if byte != b'\r' { + line_has_data = true; + } + } + Ok(chunk) + })); while let Some(event) = events.next().await { let event = event.map_err(|e| e.to_string())?; let payload = event.data; if payload.is_empty() { continue; } - let frame = TransportFrame::parse_json(&payload); + let frame = admission + .admit(TransportFrame::parse_json(&payload)) + .await + .map_err(|error| error.to_string())?; - if event_tx.unbounded_send(SseMessage { frame }).is_err() { + if event_tx.send(SseMessage { frame }).await.is_err() { return Err("upstream channel closed".to_string()); } } @@ -1196,7 +1395,7 @@ where } = channel; let writer = async move { while let Some(frame) = outgoing.next().await { - let text = match frame.to_json() { + let text = match frame.frame().to_json() { Ok(text) => text, Err(error) => { error!("failed to serialize outbound frame: {error}"); @@ -1222,7 +1421,7 @@ where continue; } let frame = TransportFrame::parse_json(text.as_str()); - if incoming.unbounded_send(frame).is_err() { + if incoming.send_frame(frame).await.is_err() { debug!( "upstream channel closed; discarding WS input while draining output" ); @@ -1256,6 +1455,10 @@ where } } +#[cfg(test)] +#[path = "client_admission_tests.rs"] +mod admission_tests; + #[cfg(test)] mod tests { use std::{ @@ -1287,9 +1490,7 @@ mod tests { struct PostsThenExitClient { finish: Arc, finished: Arc, - escaped_tx: futures::channel::oneshot::Sender< - futures::channel::mpsc::UnboundedSender, - >, + escaped_tx: futures::channel::oneshot::Sender, } struct InitializeThenExitClient { @@ -1299,7 +1500,7 @@ mod tests { struct QueueOutgoingThenText { text: Option, - outgoing: Option>, + outgoing: Option, } struct RecordingWsSink(mpsc::UnboundedSender); @@ -1339,6 +1540,12 @@ mod tests { } } + impl TransportFrameTestExt for agent_client_protocol::BudgetedFrame { + fn unwrap(self) -> RawJsonRpcMessage { + into_single_message(self.into_frame()).unwrap() + } + } + #[test] fn malformed_response_shapes_bypass_only_when_the_whole_frame_is_response_only() { let standalone_response = TransportFrame::parse_json( @@ -1374,12 +1581,13 @@ mod tests { reqwest::Client::new(), ); connection.set_connection_id("connection-1".to_string()); - let (incoming, _incoming_rx) = mpsc::unbounded(); + let (incoming, _incoming_rx) = Channel::duplex(); ClientState { connection, open_session_streams: HashSet::new(), pending_requests: HashMap::new(), - incoming, + pending_request_leases: HashMap::new(), + incoming: incoming.tx, } } @@ -1515,7 +1723,7 @@ mod tests { if let Some(outgoing) = self.outgoing.take() { for method in ["custom/first", "custom/second"] { outgoing - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(), )) .unwrap(); @@ -1559,7 +1767,7 @@ mod tests { })?; channel .tx - .unbounded_send(single_frame( + .send_frame(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -1567,19 +1775,28 @@ mod tests { ) .unwrap(), )) + .await .map_err(|e| { AcpError::internal_error().data(format!("send initialize: {e}")) })?; - into_single_message(channel.rx.next().await.ok_or_else(|| { - AcpError::internal_error().data("initialize response channel closed") - })?)?; + into_single_message( + channel + .rx + .next() + .await + .ok_or_else(|| { + AcpError::internal_error().data("initialize response channel closed") + })? + .into_frame(), + )?; for method in ["custom/first", "custom/second"] { channel .tx - .unbounded_send(single_frame( + .send_frame(single_frame( RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(), )) + .await .map_err(|e| { AcpError::internal_error().data(format!("send {method}: {e}")) })?; @@ -1605,7 +1822,7 @@ mod tests { let client = async move { channel .tx - .unbounded_send(single_frame( + .send_frame(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -1613,12 +1830,20 @@ mod tests { ) .unwrap(), )) + .await .map_err(|error| { AcpError::internal_error().data(format!("send initialize: {error}")) })?; - into_single_message(channel.rx.next().await.ok_or_else(|| { - AcpError::internal_error().data("initialize response channel closed") - })?)?; + into_single_message( + channel + .rx + .next() + .await + .ok_or_else(|| { + AcpError::internal_error().data("initialize response channel closed") + })? + .into_frame(), + )?; sse_started.notified().await; finished.notify_one(); @@ -1714,7 +1939,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -1732,7 +1957,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification( "$/cancel_request".to_string(), json!({ @@ -1832,7 +2057,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -1862,7 +2087,7 @@ mod tests { ]); caller .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ RawJsonRpcMessage::notification("custom/outbound-one".to_string(), json!({})) .unwrap(), @@ -1884,11 +2109,12 @@ mod tests { .await .unwrap() .unwrap(); - assert!(matches!(&inbound, TransportFrame::Batch(_))); + assert!(matches!(inbound.frame(), TransportFrame::Batch(_))); assert_eq!( - serde_json::from_str::(&inbound.to_json().unwrap()).unwrap(), + serde_json::from_str::(&inbound.frame().to_json().unwrap()).unwrap(), inbound_batch ); + drop(inbound); drop(caller); timeout(Duration::from_secs(1), transport) @@ -1997,7 +2223,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2013,7 +2239,7 @@ mod tests { caller .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([RawJsonRpcMessage::request( "session/fork".to_string(), json!({ "sessionId": "source-session" }), @@ -2045,11 +2271,13 @@ mod tests { .await .unwrap() .unwrap(); - assert!(matches!(&response, TransportFrame::Batch(_))); + assert!(matches!(response.frame(), TransportFrame::Batch(_))); assert_eq!( - serde_json::from_str::(&response.to_json().unwrap()).unwrap(), + serde_json::from_str::(&response.frame().to_json().unwrap()) + .unwrap(), response_batch ); + drop(response); let forked_stream = timeout(Duration::from_secs(1), get_rx.recv()) .await .unwrap() @@ -2129,7 +2357,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2153,7 +2381,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "custom/sessionish".to_string(), json!({}), @@ -2258,7 +2486,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2282,7 +2510,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "session/fork".to_string(), json!({ "sessionId": "source-session" }), @@ -2381,7 +2609,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2397,7 +2625,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification("custom/slow".to_string(), json!({})).unwrap(), )) .unwrap(); @@ -2407,7 +2635,7 @@ mod tests { caller .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ RawJsonRpcMessage::notification("custom/one".to_string(), json!({})).unwrap(), RawJsonRpcMessage::notification("custom/two".to_string(), json!({})).unwrap(), @@ -2424,7 +2652,7 @@ mod tests { caller .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ RawJsonRpcMessage::response(RequestId::Number(10), Ok(json!({}))), RawJsonRpcMessage::response(RequestId::Number(11), Ok(json!({}))), @@ -2529,7 +2757,7 @@ mod tests { ); assert!( escaped - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification("custom/too-late".to_string(), json!({}),) .unwrap() )) @@ -2622,25 +2850,30 @@ mod tests { reqwest::Client::new(), ); connection.set_connection_id("connection-1".to_string()); - let (incoming, _incoming_rx) = mpsc::unbounded(); + let (incoming, _incoming_rx) = Channel::duplex(); let mut state = ClientState { connection: connection.clone(), open_session_streams: HashSet::new(), pending_requests: HashMap::new(), - incoming, + pending_request_leases: HashMap::new(), + incoming: incoming.tx, }; let pending_request = (RequestId::Number(7), "custom/earlier".to_string()); state.track_pending_requests(std::slice::from_ref(&pending_request)); let mut posts = PostQueues::default(); posts.ordered.push(PendingPost { pending_requests: vec![pending_request], + cancelled_requests: Vec::new(), response: async { Err("earlier post failed".to_string()) }.boxed(), }); - let (_outgoing_tx, mut outgoing) = mpsc::unbounded(); + let (outgoing_channel, _outgoing_peer) = Channel::duplex(); + let mut outgoing = outgoing_channel.rx; let mut buffered_outgoing = VecDeque::new(); - let (event_tx, mut event_rx) = mpsc::unbounded(); - let mut lifecycle = HttpTransportLifecycle::new(connection); + let (event_tx, mut event_rx) = mpsc::channel(16); + let admission = state.incoming.admission(); + let max_tasks = admission.limits().max_queued_frames; + let mut lifecycle = HttpTransportLifecycle::new(connection, admission, max_tasks); let error = timeout( Duration::from_secs(1), lifecycle.start_sse( @@ -2703,16 +2936,18 @@ mod tests { reqwest::Client::new(), ); connection.set_connection_id("connection-1".to_string()); - let (incoming, mut incoming_rx) = mpsc::unbounded(); + let (incoming, mut incoming_peer) = Channel::duplex(); let mut state = ClientState { connection: connection.clone(), open_session_streams: HashSet::new(), pending_requests: HashMap::new(), - incoming, + pending_request_leases: HashMap::new(), + incoming: incoming.tx, }; let mut posts = PostQueues::default(); posts.ordered.push(PendingPost { pending_requests: Vec::new(), + cancelled_requests: Vec::new(), response: async move { complete_earlier_post.notified().await; Ok(()) @@ -2720,41 +2955,50 @@ mod tests { .boxed(), }); - let (outgoing_tx, mut outgoing) = mpsc::unbounded(); + let (outgoing_channel, outgoing_peer) = Channel::duplex(); + let outgoing_tx = outgoing_peer.tx; + let mut outgoing = outgoing_channel.rx; let outgoing_guard = outgoing_tx.clone(); let mut buffered_outgoing = VecDeque::new(); - let (event_tx, mut event_rx) = mpsc::unbounded(); + let (mut event_tx, mut event_rx) = mpsc::channel(16); event_tx - .unbounded_send(SseMessage { - frame: single_frame( - RawJsonRpcMessage::request( - "test/callback".to_string(), - json!({}), - RequestId::Number(99), - ) + .try_send(SseMessage { + frame: state + .incoming + .admission() + .try_admit(single_frame( + RawJsonRpcMessage::request( + "test/callback".to_string(), + json!({}), + RequestId::Number(99), + ) + .unwrap(), + )) .unwrap(), - ), }) .unwrap(); let responder = async move { - let callback = incoming_rx + let callback = incoming_peer + .rx .next() .await .expect("callback was not delivered"); assert!(matches!( - into_single_message(callback).unwrap(), + into_single_message(callback.into_frame()).unwrap(), RawJsonRpcMessage::Request(request) if request.method.as_ref() == "test/callback" )); outgoing_tx - .unbounded_send(single_frame(RawJsonRpcMessage::response( + .try_send(single_frame(RawJsonRpcMessage::response( RequestId::Number(99), Ok(json!({})), ))) .unwrap(); }; - let mut lifecycle = HttpTransportLifecycle::new(connection); + let admission = state.incoming.admission(); + let max_tasks = admission.limits().max_queued_frames; + let mut lifecycle = HttpTransportLifecycle::new(connection, admission, max_tasks); let (outcome, ()) = timeout(Duration::from_secs(1), async { futures::join!( lifecycle.start_sse( @@ -2821,7 +3065,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2840,7 +3084,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(), )) .unwrap(); @@ -2942,7 +3186,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2963,7 +3207,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "custom/slow".to_string(), json!({}), @@ -2987,7 +3231,7 @@ mod tests { caller .tx - .unbounded_send(single_frame(RawJsonRpcMessage::response( + .try_send(single_frame(RawJsonRpcMessage::response( RequestId::Number(99), Ok(json!({})), ))) @@ -3039,7 +3283,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3057,7 +3301,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "session/prompt".to_string(), json!({}), @@ -3103,7 +3347,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3158,7 +3402,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3179,7 +3423,7 @@ mod tests { .unwrap() .unwrap(); - let TransportFrame::Malformed { raw, error } = frame else { + let TransportFrame::Malformed { raw, error } = frame.frame() else { panic!("expected malformed frame, got {frame:?}"); }; assert_eq!(raw, "{not json"); @@ -3211,7 +3455,7 @@ mod tests { .await .unwrap() .unwrap(); - let TransportFrame::Malformed { raw, error } = frame else { + let TransportFrame::Malformed { raw, error } = frame.frame() else { panic!("expected malformed frame, got {frame:?}"); }; assert_eq!(raw, "{not json"); @@ -3243,7 +3487,7 @@ mod tests { } = caller; drop(incoming); outgoing - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ RawJsonRpcMessage::notification("custom/first".to_string(), json!({})).unwrap(), RawJsonRpcMessage::notification("custom/second".to_string(), json!({})) @@ -3335,7 +3579,7 @@ mod tests { } = caller; drop(incoming); outgoing - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(), )) .unwrap(); @@ -3416,7 +3660,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3474,7 +3718,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3525,7 +3769,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3584,7 +3828,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), diff --git a/src/agent-client-protocol-http/src/client_admission_tests.rs b/src/agent-client-protocol-http/src/client_admission_tests.rs new file mode 100644 index 00000000..9e526fcd --- /dev/null +++ b/src/agent-client-protocol-http/src/client_admission_tests.rs @@ -0,0 +1,297 @@ +use std::{convert::Infallible, time::Duration}; + +use agent_client_protocol::ConnectionLimits; +use axum::{ + Router, + response::{Sse, sse::Event}, + routing::{delete, get}, +}; +use futures::{StreamExt, channel::mpsc}; +use tokio::{net::TcpListener, time::timeout}; + +use super::*; + +#[tokio::test] +async fn sse_staging_holds_shared_budget_until_frame_is_released() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".to_string(), serde_json::json!({})).unwrap(), + ); + let json = frame.to_json().unwrap(); + let frame_bytes = json.len(); + let limits = ConnectionLimits { + max_frame_bytes: frame_bytes + 128, + // Leave room for one data event and reserve a whole frame for control. + max_queued_bytes: frame_bytes + frame_bytes + 128, + max_queued_frames: 4, + }; + let (_caller, transport) = Channel::duplex_with_limits(limits); + let admission = transport.tx.admission(); + let app = Router::new().route( + "/acp", + get({ + let json = json.clone(); + move || { + let json = json.clone(); + async move { + Sse::new(futures::stream::iter((0..3).map(move |_| { + Ok::<_, Infallible>(Event::default().data(json.clone())) + }))) + } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let connection = HttpConnection::new( + url::Url::parse(&format!("http://{address}/acp")).unwrap(), + reqwest::Client::new(), + ); + connection.set_connection_id("test-connection".into()); + let (event_tx, mut event_rx) = mpsc::channel(4); + let (established_tx, established_rx) = futures::channel::oneshot::channel(); + let reader = tokio::spawn(read_sse( + connection, + None, + event_tx, + established_tx, + admission.clone(), + )); + timeout(Duration::from_secs(2), established_rx) + .await + .unwrap() + .unwrap(); + let first = timeout(Duration::from_secs(2), event_rx.next()) + .await + .unwrap() + .unwrap(); + assert_eq!(first.frame.frame().to_json().unwrap(), json); + assert!(admission.try_admit(frame.clone()).is_err()); + // A second event may be parsed, but cannot enter the staging queue until + // the first event's shared charge is released. + assert!( + timeout(Duration::from_millis(40), event_rx.next()) + .await + .is_err() + ); + drop(first); + let second = timeout(Duration::from_secs(2), event_rx.next()) + .await + .unwrap() + .unwrap(); + assert!(admission.try_admit(frame.clone()).is_err()); + reader.abort(); + // Dropping the SSE reader releases even an event admitted but still + // waiting to send; dropping the receiver releases queued events too. + drop(second); + drop(event_rx); + reader.await.unwrap_err(); + let recovered = admission + .try_admit(frame) + .expect("cancelled SSE released permits"); + drop(recovered); + server.abort(); +} + +#[tokio::test] +async fn post_and_stream_counts_are_bounded_independently_of_frame_bytes() { + let cancellation = TransportFrame::Single( + RawJsonRpcMessage::notification( + "$/cancel_request".to_string(), + serde_json::json!({"requestId": 1}), + ) + .unwrap(), + ); + assert!(is_cancellation_frame(&cancellation)); + assert!(!is_response_only_frame(&cancellation)); + let (_caller, transport) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 4096, + max_queued_frames: 3, + }); + let connection = HttpConnection::new( + url::Url::parse("http://127.0.0.1:1/acp").unwrap(), + reqwest::Client::new(), + ); + let mut lifecycle = HttpTransportLifecycle::new(connection, transport.tx.admission(), 3); + for index in 0..3 { + drop( + lifecycle + .begin_sse(Some(index.to_string()), mpsc::channel(1).0) + .unwrap(), + ); + } + assert!( + lifecycle + .begin_sse(Some("excess".into()), mpsc::channel(1).0) + .is_err() + ); + lifecycle.sse_tasks.abort_all(); + assert_eq!(lifecycle.sse_tasks.len(), 0); + + let mut posts = PostQueues::default(); + for _ in 0..2 { + check_post_capacity(&posts, 3, false).unwrap(); + posts.ordered.push(PendingPost { + pending_requests: Vec::new(), + cancelled_requests: Vec::new(), + response: Box::pin(futures::future::pending()), + }); + } + assert!(check_post_capacity(&posts, 3, false).is_err()); + check_post_capacity(&posts, 3, true).unwrap(); + posts.responses.push(PendingPost { + pending_requests: Vec::new(), + cancelled_requests: Vec::new(), + response: Box::pin(futures::future::pending()), + }); + assert_eq!(posts.len(), 3); + assert!(check_post_capacity(&posts, 3, true).is_err()); + drop(posts); +} + +#[tokio::test] +async fn cancelled_post_releases_its_body_budget() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".to_string(), serde_json::json!({})).unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let (_caller, transport) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes + 128, + max_queued_bytes: bytes * 2 + 128, + max_queued_frames: 4, + }); + let admission = transport.tx.admission(); + let budgeted = admission.try_admit(frame.clone()).unwrap(); + let (_, permit) = budgeted.into_parts(); + let mut posts = PostQueues::default(); + posts.ordered.push_budgeted( + PendingPost { + pending_requests: Vec::new(), + cancelled_requests: Vec::new(), + response: Box::pin(futures::future::pending()), + }, + permit, + ); + assert!(admission.try_admit(frame.clone()).is_err()); + drop(posts); + assert!(admission.try_admit(frame).is_ok()); +} + +#[tokio::test] +async fn delivering_sse_preserves_admission_through_output_channel() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".to_string(), serde_json::json!({})).unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let (mut caller, transport) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes + 128, + max_queued_bytes: bytes * 2 + 128, + max_queued_frames: 4, + }); + let admission = transport.tx.admission(); + let state = ClientState { + connection: HttpConnection::new( + url::Url::parse("http://127.0.0.1:1/acp").unwrap(), + reqwest::Client::new(), + ), + open_session_streams: HashSet::new(), + pending_requests: HashMap::new(), + pending_request_leases: HashMap::new(), + incoming: transport.tx, + }; + state + .deliver_budgeted(admission.try_admit(frame.clone()).unwrap()) + .await + .unwrap(); + assert!(admission.try_admit(frame.clone()).is_err()); + let delivered = caller.rx.next().await.unwrap(); + assert!(admission.try_admit(frame.clone()).is_err()); + drop(delivered); + assert!(admission.try_admit(frame).is_ok()); +} + +#[tokio::test] +async fn pending_requests_hold_their_source_charge_until_response_or_cancel() { + let request = RawJsonRpcMessage::request( + "test/request".to_string(), + serde_json::json!({}), + RequestId::Number(1), + ) + .unwrap(); + let frame = TransportFrame::Single(request); + let bytes = frame.to_json().unwrap().len(); + let (_caller, transport) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes + 128, + max_queued_bytes: bytes * 2 + 128, + max_queued_frames: 1, + }); + let admission = transport.tx.admission(); + let mut state = ClientState { + connection: HttpConnection::new( + url::Url::parse("http://127.0.0.1:1/acp").unwrap(), + reqwest::Client::new(), + ), + open_session_streams: HashSet::new(), + pending_requests: HashMap::new(), + pending_request_leases: HashMap::new(), + incoming: transport.tx, + }; + state.connection.set_connection_id("connection-1".into()); + let (_, permit) = admission.try_admit(frame.clone()).unwrap().into_parts(); + let post = state.prepare_frame_post(frame.clone()).unwrap().0; + state.attach_pending_permits(&post.pending_requests, &permit); + assert!(state.check_pending_request_capacity(1).is_err()); + drop(post); + drop(permit); + assert!(admission.try_admit(frame.clone()).is_err()); + assert_eq!( + state + .take_pending_request_method(&RequestId::Number(1)) + .as_deref(), + Some("test/request") + ); + assert!(state.check_pending_request_capacity(1).is_ok()); + let (_, permit) = admission.try_admit(frame.clone()).unwrap().into_parts(); + let post = state.prepare_frame_post(frame.clone()).unwrap().0; + state.attach_pending_permits(&post.pending_requests, &permit); + drop(post); + drop(permit); + let cancel = RawJsonRpcMessage::notification( + "$/cancel_request".into(), + serde_json::json!({"requestId": 1}), + ) + .unwrap(); + let post = state.prepare_post(cancel).unwrap(); + handle_completed_post( + &mut state, + CompletedPost { + pending_requests: post.pending_requests, + cancelled_requests: post.cancelled_requests, + result: Ok(()), + }, + ) + .unwrap(); + assert!(state.check_pending_request_capacity(1).is_ok()); + assert!(admission.try_admit(frame).is_ok()); +} + +#[tokio::test] +async fn unresponsive_close_does_not_stall_transport_shutdown() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let app = Router::new().route( + "/acp", + delete(|| async { futures::future::pending::().await }), + ); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let connection = HttpConnection::new( + url::Url::parse(&format!("http://{address}/acp")).unwrap(), + reqwest::Client::new(), + ); + connection.set_connection_id("test-connection".into()); + timeout(Duration::from_secs(4), connection.close()) + .await + .expect("DELETE must not indefinitely block transport shutdown"); + server.abort(); +} diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index 3fdbe0f5..23b5591b 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -4,8 +4,9 @@ use std::{ }; use agent_client_protocol::{ - Channel, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame, - schema::v1::RequestId, + BudgetedFrame, Channel, ConnectionLimits, FrameAdmission, FramePermit, RawJsonRpcMessage, + TransportBatch, TransportBatchEntry, TransportFrame, + schema::v1::{RequestId, Response as RpcResponse}, }; use futures::{SinkExt, StreamExt}; use tokio::sync::{Mutex, RwLock, mpsc, watch}; @@ -20,14 +21,15 @@ pub(crate) enum ResponseRoute { } enum OutboundTransport { - Http(HttpOutbound), + Http(Box), WebSocket(WebSocketOutbound), } struct HttpOutbound { connection_stream: OutboundMailbox, - session_streams: RwLock>>, - pending_routes: Mutex>>, + session_streams: RwLock, Option)>>, + pending_routes: Mutex)>>>, + limits: ConnectionLimits, } struct WebSocketOutbound { @@ -35,28 +37,43 @@ struct WebSocketOutbound { } struct OutboundMailbox { - sender: mpsc::UnboundedSender, - receiver_slot: Arc>>>, + sender: mpsc::Sender, + receiver_slot: Arc>>>, +} + +struct OutboundValue { + text: String, + permit: Option, } pub(crate) struct OutboundLease { - receiver: Option>, - receiver_slot: Arc>>>, + receiver: Option>, + receiver_slot: Arc>>>, + current: Option, } impl OutboundMailbox { fn new() -> Self { - let (sender, receiver) = mpsc::unbounded_channel(); + let (sender, receiver) = mpsc::channel(32); Self { sender, receiver_slot: Arc::new(StdMutex::new(Some(receiver))), } } + #[cfg(test)] fn push(&self, msg: String) -> Result<(), &'static str> { + self.push_with_permit(msg, None) + } + + fn push_with_permit( + &self, + text: String, + permit: Option, + ) -> Result<(), &'static str> { self.sender - .send(msg) - .map_err(|_| "outbound mailbox receiver closed") + .try_send(OutboundValue { text, permit }) + .map_err(|_| "outbound mailbox full or receiver closed") } fn try_acquire(&self) -> Option { @@ -68,24 +85,31 @@ impl OutboundMailbox { Some(OutboundLease { receiver: Some(receiver), receiver_slot: self.receiver_slot.clone(), + current: None, }) } } impl OutboundLease { pub(crate) async fn recv(&mut self) -> Option { - self.receiver + let value = self + .receiver .as_mut() .expect("outbound lease receiver missing") .recv() - .await + .await?; + self.current = value.permit; + Some(value.text) } pub(crate) fn try_recv(&mut self) -> Result { - self.receiver + let value = self + .receiver .as_mut() .expect("outbound lease receiver missing") - .try_recv() + .try_recv()?; + self.current = value.permit; + Ok(value.text) } } @@ -104,8 +128,9 @@ impl Drop for OutboundLease { } pub(crate) struct Connection { - inbound_tx: mpsc::UnboundedSender, - outbound_rx: Mutex>>, + inbound_tx: mpsc::Sender, + inbound_admission: FrameAdmission, + outbound_rx: Mutex>>, agent_handle: Mutex>>, router_handle: Mutex>>, closed_tx: watch::Sender, @@ -114,17 +139,61 @@ pub(crate) struct Connection { impl Connection { pub(crate) fn send_frame_to_agent(&self, frame: TransportFrame) -> Result<(), &'static str> { + let frame = self.admit_frame_to_agent(frame)?; + self.send_budgeted_frame_to_agent(frame) + } + + pub(crate) fn admit_frame_to_agent( + &self, + frame: TransportFrame, + ) -> Result { + self.inbound_admission + .try_admit(frame) + .map_err(|_| "agent frame byte capacity exceeded") + } + + pub(crate) fn send_budgeted_frame_to_agent( + &self, + frame: BudgetedFrame, + ) -> Result<(), &'static str> { self.inbound_tx - .send(frame) - .map_err(|_| "agent channel closed") + .try_send(frame) + .map_err(|_| "agent channel full or closed") } - pub(crate) async fn record_pending_route(&self, id: RequestId, route: ResponseRoute) { - self.outbound_transport - .record_pending_route(id, route) - .await; + pub(crate) async fn register_post_routes( + &self, + sessions: &[String], + routes: &[(RequestId, ResponseRoute)], + permit: &FramePermit, + ) -> Result, &'static str> { + if let OutboundTransport::Http(http) = &self.outbound_transport { + http.register_post_routes(sessions, routes, permit).await + } else { + Ok(Vec::new()) + } + } + + pub(crate) async fn rollback_post_routes( + &self, + sessions: &[String], + routes: &[(RequestId, ResponseRoute)], + ) { + if let OutboundTransport::Http(http) = &self.outbound_transport { + http.rollback_post_routes(sessions, routes).await; + } } + pub(crate) async fn cancel_pending_routes(&self, ids: &[RequestId]) { + if let OutboundTransport::Http(http) = &self.outbound_transport { + let mut pending = http.pending_routes.lock().await; + for id in ids { + take_pending_route(&mut pending, id); + } + } + } + + #[cfg(test)] pub(crate) async fn ensure_session(&self, session_id: &str) { self.outbound_transport.ensure_session(session_id).await; } @@ -177,11 +246,14 @@ impl Connection { })); } - pub(crate) async fn route_outbound(&self, frame: TransportFrame) -> Result<(), &'static str> { - self.outbound_transport.route_outbound(frame).await + pub(crate) async fn route_outbound(&self, frame: BudgetedFrame) -> Result<(), &'static str> { + let (frame, permit) = frame.into_parts(); + self.outbound_transport + .route_outbound(frame, Some(permit)) + .await } - pub(crate) async fn recv_initial(&self) -> Option { + pub(crate) async fn recv_initial(&self) -> Option { let mut guard = self.outbound_rx.lock().await; let rx = guard.as_mut()?; rx.recv().await @@ -191,6 +263,10 @@ impl Connection { // Explicit peer teardown is abortive. Natural agent completion instead // awaits the router in `close_connection_task` before closing streams. self.close_streams(); + if let OutboundTransport::Http(http) = &self.outbound_transport { + http.session_streams.write().await.clear(); + http.pending_routes.lock().await.clear(); + } if let Some(h) = self.agent_handle.lock().await.take() { h.abort(); } @@ -206,21 +282,14 @@ impl Connection { impl OutboundTransport { fn http() -> Self { - Self::Http(HttpOutbound::new()) + Self::Http(Box::new(HttpOutbound::new())) } fn websocket() -> Self { Self::WebSocket(WebSocketOutbound::new()) } - async fn record_pending_route(&self, id: RequestId, route: ResponseRoute) { - let Self::Http(http) = self else { - return; - }; - - http.record_pending_route(id, route).await; - } - + #[cfg(test)] async fn ensure_session(&self, session_id: &str) { let Self::Http(http) = self else { return; @@ -238,7 +307,12 @@ impl OutboundTransport { async fn subscribe_session_stream(&self, session_id: &str) -> Option { match self { - Self::Http(http) => http.session_stream(session_id).await.try_acquire(), + Self::Http(http) => http + .session_streams + .read() + .await + .get(session_id) + .and_then(|(stream, _)| stream.try_acquire()), Self::WebSocket(_) => None, } } @@ -259,7 +333,11 @@ impl OutboundTransport { http.connection_stream.push(msg) } - async fn route_outbound(&self, frame: TransportFrame) -> Result<(), &'static str> { + async fn route_outbound( + &self, + frame: TransportFrame, + permit: Option, + ) -> Result<(), &'static str> { match frame { TransportFrame::Single(message) => { let serialized = match serde_json::to_string(&message) { @@ -270,13 +348,18 @@ impl OutboundTransport { } }; match self { - Self::Http(http) => http.route_outbound(&message, serialized).await, - Self::WebSocket(websocket) => websocket.all_outbound.push(serialized), + Self::Http(http) => { + http.route_outbound_with_permit(&message, serialized, permit) + .await + } + Self::WebSocket(websocket) => { + websocket.all_outbound.push_with_permit(serialized, permit) + } } } TransportFrame::Malformed { raw, .. } => match self { - Self::Http(http) => http.connection_stream.push(raw), - Self::WebSocket(websocket) => websocket.all_outbound.push(raw), + Self::Http(http) => http.connection_stream.push_with_permit(raw, permit), + Self::WebSocket(websocket) => websocket.all_outbound.push_with_permit(raw, permit), }, TransportFrame::Batch(batch) => { let serialized = match serde_json::to_string(&batch) { @@ -287,8 +370,10 @@ impl OutboundTransport { } }; match self { - Self::Http(http) => http.route_outbound_batch(&batch, serialized).await, - Self::WebSocket(websocket) => websocket.all_outbound.push(serialized), + Self::Http(http) => http.route_outbound_batch(&batch, serialized, permit).await, + Self::WebSocket(websocket) => { + websocket.all_outbound.push_with_permit(serialized, permit) + } } } } @@ -301,9 +386,70 @@ impl HttpOutbound { connection_stream: OutboundMailbox::new(), session_streams: RwLock::new(HashMap::new()), pending_routes: Mutex::new(HashMap::new()), + limits: ConnectionLimits::default(), } } + async fn register_post_routes( + &self, + sessions: &[String], + routes: &[(RequestId, ResponseRoute)], + permit: &FramePermit, + ) -> Result, &'static str> { + // Lock both metadata tables in one order and check the whole batch + // before inserting anything: rejection must never leave half a batch. + let mut streams = self.session_streams.write().await; + let mut pending = self.pending_routes.lock().await; + let mut new_sessions = Vec::new(); + for id in sessions { + if !streams.contains_key(id) && !new_sessions.contains(id) { + new_sessions.push(id.clone()); + } + } + let pending_count: usize = pending.values().map(VecDeque::len).sum(); + let limit = self.limits.max_queued_frames.max(1); + let available = limit.saturating_sub(streams.len().saturating_add(pending_count)); + if new_sessions.len().saturating_add(routes.len()) > available { + return Err("HTTP pending route or session capacity exceeded"); + } + for id in &new_sessions { + streams.insert( + id.clone(), + (Arc::new(OutboundMailbox::new()), Some(permit.clone())), + ); + } + for (id, route) in routes { + if let Some(id) = pending_route_key(id) { + pending + .entry(id) + .or_default() + .push_back((route.clone(), Some(permit.clone()))); + } + } + Ok(new_sessions) + } + + async fn rollback_post_routes( + &self, + new_sessions: &[String], + routes: &[(RequestId, ResponseRoute)], + ) { + let mut streams = self.session_streams.write().await; + let mut pending = self.pending_routes.lock().await; + for id in new_sessions { + streams.remove(id); + } + for (id, _) in routes.iter().rev() { + if let Some(queue) = pending.get_mut(id) { + queue.pop_back(); + if queue.is_empty() { + pending.remove(id); + } + } + } + } + + #[cfg(test)] async fn record_pending_route(&self, id: RequestId, route: ResponseRoute) { if let Some(key) = pending_route_key(&id) { self.pending_routes @@ -311,31 +457,71 @@ impl HttpOutbound { .await .entry(key) .or_default() - .push_back(route); + .push_back((route, None)); } } + #[cfg(test)] async fn ensure_session(&self, session_id: &str) { self.session_stream(session_id).await; } + async fn session_stream_with_permit( + &self, + session_id: &str, + permit: Option, + ) -> Result, &'static str> { + let mut streams = self.session_streams.write().await; + if let Some((stream, _)) = streams.get(session_id) { + return Ok(stream.clone()); + } + let Some(permit) = permit else { + return Err("session stream has no admitted source frame"); + }; + let pending_count: usize = self + .pending_routes + .lock() + .await + .values() + .map(VecDeque::len) + .sum(); + if streams.len().saturating_add(pending_count) >= self.limits.max_queued_frames.max(1) { + return Err("HTTP session stream capacity exceeded"); + } + let stream = Arc::new(OutboundMailbox::new()); + streams.insert(session_id.to_string(), (stream.clone(), Some(permit))); + Ok(stream) + } + + #[cfg(test)] async fn session_stream(&self, session_id: &str) -> Arc { if let Some(stream) = self.session_streams.read().await.get(session_id) { - return stream.clone(); + return stream.0.clone(); } self.session_streams .write() .await .entry(session_id.to_string()) - .or_insert_with(|| Arc::new(OutboundMailbox::new())) + .or_insert_with(|| (Arc::new(OutboundMailbox::new()), None)) + .0 .clone() } + #[cfg(test)] async fn route_outbound( &self, msg: &RawJsonRpcMessage, serialized: String, + ) -> Result<(), &'static str> { + self.route_outbound_with_permit(msg, serialized, None).await + } + + async fn route_outbound_with_permit( + &self, + msg: &RawJsonRpcMessage, + serialized: String, + permit: Option, ) -> Result<(), &'static str> { let route = match msg { RawJsonRpcMessage::Request(_) | RawJsonRpcMessage::Notification(_) => { @@ -353,15 +539,23 @@ impl HttpOutbound { route.unwrap_or(ResponseRoute::Connection) } }; + // A successful session/new (or fork) response can be followed + // immediately by a session SSE GET, before any session-scoped POST. + if let Some(session_id) = response_session_id(msg) { + self.session_stream_with_permit(session_id, permit.clone()) + .await?; + } match route { ResponseRoute::Connection => { trace!(target = "connection", "→ connection-scoped stream"); - self.connection_stream.push(serialized) + self.connection_stream.push_with_permit(serialized, permit) } ResponseRoute::Session(sid) => { trace!(target = %sid, "→ session-scoped stream"); - self.session_stream(&sid).await.push(serialized) + self.session_stream_with_permit(&sid, permit.clone()) + .await? + .push_with_permit(serialized, permit) } } } @@ -370,6 +564,7 @@ impl HttpOutbound { &self, batch: &TransportBatch, serialized: String, + permit: Option, ) -> Result<(), &'static str> { let mut pending_routes = self.pending_routes.lock().await; let mut common_route = None; @@ -390,6 +585,14 @@ impl HttpOutbound { } } drop(pending_routes); + for entry in batch.entries() { + if let TransportBatchEntry::Message(message) = entry + && let Some(session_id) = response_session_id(message) + { + self.session_stream_with_permit(session_id, permit.clone()) + .await?; + } + } let route = if routes_disagree { ResponseRoute::Connection @@ -399,11 +602,13 @@ impl HttpOutbound { match route { ResponseRoute::Connection => { trace!(target = "connection", "→ connection-scoped batch stream"); - self.connection_stream.push(serialized) + self.connection_stream.push_with_permit(serialized, permit) } ResponseRoute::Session(session_id) => { trace!(target = %session_id, "→ session-scoped batch stream"); - self.session_stream(&session_id).await.push(serialized) + self.session_stream_with_permit(&session_id, permit.clone()) + .await? + .push_with_permit(serialized, permit) } } } @@ -442,7 +647,8 @@ where Channel, futures::future::BoxFuture<'static, agent_client_protocol::Result<()>>, ) { - self().into_channel_and_future() + let (channel, driver) = self().into_channel_and_future(); + (channel, Box::pin(driver)) } } @@ -483,14 +689,19 @@ impl ConnectionRegistry { outbound_transport: OutboundTransport, ) -> Arc { let (channel, agent_future) = self.factory.spawn_agent(); - let (inbound_tx, mut inbound_rx) = mpsc::unbounded_channel::(); - let (outbound_tx, outbound_rx) = mpsc::unbounded_channel::(); + let mut outbound_transport = outbound_transport; + if let OutboundTransport::Http(http) = &mut outbound_transport { + http.limits = channel.tx.admission().limits(); + } + let (inbound_tx, mut inbound_rx) = mpsc::channel::(32); + let (outbound_tx, outbound_rx) = mpsc::channel::(32); let (closed_tx, _) = watch::channel(false); let Channel { rx: mut agent_rx, tx: mut agent_tx, } = channel; + let inbound_admission = agent_tx.admission(); let inbound = async move { while let Some(msg) = inbound_rx.recv().await { if agent_tx.send(msg).await.is_err() { @@ -504,7 +715,7 @@ impl ConnectionRegistry { let inbound_abort_for_outbound = inbound_abort.clone(); let outbound = async move { while let Some(msg) = agent_rx.next().await { - if outbound_tx.send(msg).is_err() { + if outbound_tx.send(msg).await.is_err() { inbound_abort_for_outbound.abort(); break; } @@ -516,6 +727,7 @@ impl ConnectionRegistry { let connection = Arc::new(Connection { inbound_tx, + inbound_admission, outbound_rx: Mutex::new(Some(outbound_rx)), agent_handle: Mutex::new(None), router_handle: Mutex::new(None), @@ -582,6 +794,10 @@ async fn close_connection_task(connection: Weak) { error!("outbound router task failed while draining: {error}"); } connection.close_streams(); + if let OutboundTransport::Http(http) = &connection.outbound_transport { + http.session_streams.write().await.clear(); + http.pending_routes.lock().await.clear(); + } } fn pending_route_key(id: &RequestId) -> Option { @@ -591,8 +807,15 @@ fn pending_route_key(id: &RequestId) -> Option { } } +fn response_session_id(msg: &RawJsonRpcMessage) -> Option<&str> { + let RawJsonRpcMessage::Response(RpcResponse::Result { result, .. }) = msg else { + return None; + }; + result.get("sessionId")?.as_str() +} + fn take_pending_route( - pending_routes: &mut HashMap>, + pending_routes: &mut HashMap)>>, key: &RequestId, ) -> Option { let routes = pending_routes.get_mut(key)?; @@ -601,9 +824,13 @@ fn take_pending_route( if remove_entry { pending_routes.remove(key); } - route + route.map(|(route, _permit)| route) } +#[cfg(test)] +#[path = "connection_admission_tests.rs"] +mod admission_tests; + #[cfg(test)] mod tests { use std::sync::Arc; @@ -617,18 +844,18 @@ mod tests { use super::*; - const ISSUE_288_BURST: usize = 1_025; - #[tokio::test] - async fn outbound_mailbox_buffers_bursts_before_subscription() { + async fn outbound_mailbox_bounds_bursts_before_subscription() { let mailbox = OutboundMailbox::new(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { mailbox.push(format!("message-{index}")).unwrap(); } + assert!(mailbox.push("overflow".into()).is_err()); let mut receiver = mailbox.try_acquire().unwrap(); - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { assert_eq!( receiver.recv().await, Some(format!("message-{index}")), @@ -641,17 +868,21 @@ mod tests { async fn outbound_mailbox_does_not_stall_when_subscriber_is_slow() { let mailbox = OutboundMailbox::new(); let mut receiver = mailbox.try_acquire().unwrap(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { mailbox.push(format!("message-{index}")).unwrap(); } - for index in 0..ISSUE_288_BURST { + assert!(mailbox.push("overflow".into()).is_err()); + for index in 0..capacity { assert_eq!( receiver.recv().await, Some(format!("message-{index}")), "message {index} should remain ordered" ); } + mailbox.push("recovered".into()).unwrap(); + assert_eq!(receiver.recv().await.as_deref(), Some("recovered")); } #[tokio::test] @@ -679,6 +910,7 @@ mod tests { #[tokio::test] async fn slow_session_mailbox_does_not_stall_other_routes() { let outbound = HttpOutbound::new(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; let mut slow_session = outbound .session_stream("slow-session") .await @@ -691,7 +923,7 @@ mod tests { .unwrap(); timeout(Duration::from_secs(1), async { - for index in 0..ISSUE_288_BURST { + for index in 0..=capacity { let message = RawJsonRpcMessage::notification( "session/update".to_string(), serde_json::json!({ @@ -701,7 +933,15 @@ mod tests { ) .unwrap(); let serialized = serde_json::to_string(&message).unwrap(); - outbound.route_outbound(&message, serialized).await.unwrap(); + let result = outbound.route_outbound(&message, serialized).await; + if index == capacity { + assert!( + result.is_err(), + "overflow must be explicit, not silently dropped" + ); + } else { + result.unwrap(); + } } let marker = RawJsonRpcMessage::notification( @@ -727,7 +967,7 @@ mod tests { true ); - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { let message = slow_session.recv().await.unwrap(); assert_eq!( serde_json::from_str::(&message).unwrap()["params"]["index"], @@ -772,10 +1012,11 @@ mod tests { let future = Box::pin(async move { agent .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( RequestId::Number(1), Ok(serde_json::json!({ "done": true })), ))) + .await .unwrap(); Ok(()) }); @@ -801,11 +1042,12 @@ mod tests { emit.notified().await; agent .tx - .unbounded_send(TransportFrame::Malformed { + .send_frame(TransportFrame::Malformed { raw: "{not json".to_string(), error: agent_client_protocol::Error::parse_error() .data("transport parse error"), }) + .await .unwrap(); std::future::pending::>().await }); @@ -832,7 +1074,8 @@ mod tests { let future = Box::pin(async move { agent .tx - .unbounded_send(TransportFrame::Single(message)) + .send_frame(TransportFrame::Single(message)) + .await .unwrap(); exit.notified().await; Ok(()) @@ -871,7 +1114,8 @@ mod tests { .expect("test batch is non-empty"); agent .tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .unwrap(); exit.notified().await; Ok(()) @@ -898,13 +1142,14 @@ mod tests { emit.notified().await; agent .tx - .unbounded_send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::notification( "test/final".to_string(), serde_json::json!({}), ) .unwrap(), )) + .await .unwrap(); Ok(()) }); @@ -974,7 +1219,7 @@ mod tests { .expect("buffered response should be forwarded before teardown"); assert!(matches!( - frame, + frame.frame(), TransportFrame::Single(RawJsonRpcMessage::Response( agent_client_protocol::schema::v1::Response::Result { id: RequestId::Number(1), @@ -1045,6 +1290,7 @@ mod tests { })); let (_connection_id, connection) = registry.create_connection().await; let mut connection_rx = connection.subscribe_connection_stream().unwrap(); + connection.ensure_session("session-1").await; let mut session_rx = connection .subscribe_session_stream("session-1") .await @@ -1118,7 +1364,7 @@ mod tests { let serialized = serde_json::to_string(&batch).unwrap(); outbound - .route_outbound_batch(&batch, serialized.clone()) + .route_outbound_batch(&batch, serialized.clone(), None) .await .unwrap(); diff --git a/src/agent-client-protocol-http/src/connection_admission_tests.rs b/src/agent-client-protocol-http/src/connection_admission_tests.rs new file mode 100644 index 00000000..23066f94 --- /dev/null +++ b/src/agent-client-protocol-http/src/connection_admission_tests.rs @@ -0,0 +1,101 @@ +use agent_client_protocol::ConnectionLimits; +use serde_json::json; + +use super::*; + +#[tokio::test] +async fn route_and_session_metadata_admission_is_atomic_and_releases_permits() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::request("test/request".into(), json!({}), RequestId::Number(1)).unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes + 128, + max_queued_bytes: bytes * 2 + 128, + max_queued_frames: 2, + }; + let (_caller, transport) = Channel::duplex_with_limits(limits); + let admission = transport.tx.admission(); + let (_, permit) = admission.try_admit(frame.clone()).unwrap().into_parts(); + let mut http = HttpOutbound::new(); + http.limits = limits; + assert!( + OutboundTransport::Http(Box::new(HttpOutbound::new())) + .subscribe_session_stream("unknown") + .await + .is_none() + ); + let first = [(RequestId::Number(1), ResponseRoute::Session("one".into()))]; + http.register_post_routes(&["one".into()], &first, &permit) + .await + .unwrap(); + let extra = [(RequestId::Number(2), ResponseRoute::Session("two".into()))]; + assert!( + http.register_post_routes(&["two".into()], &extra, &permit) + .await + .is_err() + ); + assert_eq!(http.session_streams.read().await.len(), 1); + assert_eq!(http.pending_routes.lock().await.len(), 1); + drop(permit); + assert!(admission.try_admit(frame.clone()).is_err()); + assert_eq!( + take_pending_route( + &mut *http.pending_routes.lock().await, + &RequestId::Number(1) + ), + Some(ResponseRoute::Session("one".into())) + ); + assert!(admission.try_admit(frame.clone()).is_err()); + http.session_streams.write().await.clear(); + assert!(admission.try_admit(frame).is_ok()); +} + +#[tokio::test] +async fn rolling_back_rejected_transport_send_removes_only_new_metadata() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::request("test/request".into(), json!({}), RequestId::Number(1)).unwrap(), + ); + let (_caller, transport) = Channel::duplex(); + let (_, permit) = transport + .tx + .admission() + .try_admit(frame) + .unwrap() + .into_parts(); + let http = HttpOutbound::new(); + let routes = [(RequestId::Number(1), ResponseRoute::Session("one".into()))]; + let new_sessions = http + .register_post_routes(&["one".into()], &routes, &permit) + .await + .unwrap(); + http.rollback_post_routes(&new_sessions, &routes).await; + assert!(http.pending_routes.lock().await.is_empty()); + assert!(http.session_streams.read().await.is_empty()); +} + +#[tokio::test] +async fn successful_session_response_registers_stream_before_get() { + let response = RawJsonRpcMessage::response( + RequestId::Number(1), + Ok(json!({"sessionId": "new-session"})), + ); + let frame = TransportFrame::Single(response.clone()); + let (_, channel) = Channel::duplex(); + let (_, permit) = channel + .tx + .admission() + .try_admit(frame.clone()) + .unwrap() + .into_parts(); + let http = HttpOutbound::new(); + http.route_outbound_with_permit(&response, frame.to_json().unwrap(), Some(permit)) + .await + .unwrap(); + assert!( + http.session_streams + .read() + .await + .contains_key("new-session") + ); +} diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index 525b8e7c..a3d6676d 100644 --- a/src/agent-client-protocol-http/src/http_server.rs +++ b/src/agent-client-protocol-http/src/http_server.rs @@ -102,7 +102,9 @@ pub(crate) async fn handle_post( ) .into_response(); }; - if let Some(initialize_failed) = initialize_response_failed(&frame, &initialize_id) { + if let Some(initialize_failed) = + initialize_response_failed(frame.frame(), &initialize_id) + { break (frame, initialize_failed); } @@ -114,7 +116,7 @@ pub(crate) async fn handle_post( return (StatusCode::INTERNAL_SERVER_ERROR, error).into_response(); } }; - let init_response = match init_response_frame.to_json() { + let init_response = match init_response_frame.frame().to_json() { Ok(response) => response, Err(e) => { initialize_cleanup.cleanup().await; @@ -143,6 +145,7 @@ pub(crate) async fn handle_post( let mut session_routes = Vec::new(); let mut pending_routes = Vec::new(); + let mut cancellations = Vec::new(); match &mut frame { TransportFrame::Single(message) => { let route = match prepare_message_route(message, session_id.as_deref()) { @@ -150,6 +153,7 @@ pub(crate) async fn handle_post( Err(error) => return (StatusCode::BAD_REQUEST, error).into_response(), }; collect_route(message, route, &mut session_routes, &mut pending_routes); + cancellations.extend(crate::protocol::cancelled_request_id(message)); trace!(connection_id = %connection_id, ?message, "POST → agent"); } TransportFrame::Batch(batch) => { @@ -162,6 +166,7 @@ pub(crate) async fn handle_post( Err(error) => return (StatusCode::BAD_REQUEST, error).into_response(), }; collect_route(message, route, &mut session_routes, &mut pending_routes); + cancellations.extend(crate::protocol::cancelled_request_id(message)); } trace!(connection_id = %connection_id, ?frame, "POST batch → agent"); } @@ -170,16 +175,26 @@ pub(crate) async fn handle_post( } } - for session_id in session_routes { - connection.ensure_session(&session_id).await; - } - for (request_id, route) in pending_routes { - connection.record_pending_route(request_id, route).await; - } - - if connection.send_frame_to_agent(frame).is_err() { + let admitted = match connection.admit_frame_to_agent(frame) { + Ok(frame) => frame, + Err(error) => return (StatusCode::TOO_MANY_REQUESTS, error).into_response(), + }; + let permit = admitted.permit().clone(); + let new_sessions = match connection + .register_post_routes(&session_routes, &pending_routes, &permit) + .await + { + Ok(new_sessions) => new_sessions, + Err(error) => return (StatusCode::TOO_MANY_REQUESTS, error).into_response(), + }; + drop(permit); + if connection.send_budgeted_frame_to_agent(admitted).is_err() { + connection + .rollback_post_routes(&new_sessions, &pending_routes) + .await; return StatusCode::INTERNAL_SERVER_ERROR.into_response(); } + connection.cancel_pending_routes(&cancellations).await; StatusCode::ACCEPTED.into_response() } @@ -356,7 +371,7 @@ pub(crate) async fn handle_get( let Some(mut receiver) = receiver else { return ( StatusCode::CONFLICT, - "outbound stream already has a subscriber", + "outbound stream missing or already has a subscriber", ) .into_response(); }; @@ -480,8 +495,8 @@ mod tests { use std::sync::Arc; use agent_client_protocol::{ - Channel, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame, - schema::v1::RequestId, + BudgetedFrame, Channel, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, + TransportFrame, schema::v1::RequestId, }; use futures::{StreamExt, future::BoxFuture}; use serde_json::json; @@ -493,8 +508,6 @@ mod tests { use super::*; use crate::connection::AgentFactory; - const ISSUE_288_BURST: usize = 1_025; - struct CapturingAgentFactory { forwarded: mpsc::UnboundedSender, } @@ -514,7 +527,7 @@ mod tests { tx: _, } = agent; while let Some(frame) = incoming.next().await { - let TransportFrame::Single(message) = frame else { + let TransportFrame::Single(message) = frame.into_frame() else { panic!("expected a single JSON-RPC frame"); }; if forwarded.send(message).is_err() { @@ -539,15 +552,16 @@ mod tests { ) { let (mut agent, transport) = Channel::duplex(); let future = Box::pin(async move { - match agent.rx.next().await { + match agent.rx.next().await.map(BudgetedFrame::into_frame) { Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) => { agent .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Err(agent_client_protocol::Error::invalid_request() .data("initialize rejected")), ))) + .await .unwrap(); } Some(TransportFrame::Batch(batch)) => { @@ -567,10 +581,11 @@ mod tests { }); agent .tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages(responses) .expect("request batch has responses"), )) + .await .unwrap(); } Some(TransportFrame::Single(_) | TransportFrame::Malformed { .. }) | None => {} @@ -619,7 +634,9 @@ mod tests { let (mut agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); let future = Box::pin(async move { - let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { + let Some(TransportFrame::Batch(batch)) = + agent.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected one batch frame"); }; let mut methods = Vec::new(); @@ -666,7 +683,8 @@ mod tests { TransportBatch::from_messages(responses).expect("responses are non-empty"); agent .tx - .unbounded_send(TransportFrame::Batch(responses)) + .send_frame(TransportFrame::Batch(responses)) + .await .unwrap(); std::future::pending::>().await }); @@ -686,7 +704,9 @@ mod tests { ) { let (mut agent, transport) = Channel::duplex(); let future = Box::pin(async move { - let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { + let Some(TransportFrame::Batch(batch)) = + agent.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected one initial batch frame"); }; let responses = batch.entries().filter_map(|entry| { @@ -702,20 +722,22 @@ mod tests { agent .tx - .unbounded_send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::notification( "custom/during-initialize".into(), json!({ "phase": "before-response" }), ) .expect("test notification should serialize"), )) + .await .unwrap(); agent .tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages(responses) .expect("initial batch has response-bearing requests"), )) + .await .unwrap(); std::future::pending::>().await }); @@ -989,6 +1011,7 @@ mod tests { }))); let (connection_id, connection) = registry.create_connection().await; let mut connection_outbound = connection.subscribe_connection_stream().unwrap(); + connection.ensure_session("session-1").await; let mut session_outbound = connection .subscribe_session_stream("session-1") .await @@ -1057,6 +1080,7 @@ mod tests { }))); let (connection_id, connection) = registry.create_connection().await; let mut connection_outbound = connection.subscribe_connection_stream().unwrap(); + connection.ensure_session("session-1").await; let mut session_outbound = connection .subscribe_session_stream("session-1") .await @@ -1258,8 +1282,9 @@ mod tests { } #[tokio::test] - async fn sse_buffers_burst_without_polling_slow_subscriber() { + async fn sse_bounds_burst_and_drains_every_accepted_message() { let (forwarded_tx, _forwarded_rx) = mpsc::unbounded_channel(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; let registry = Arc::new(ConnectionRegistry::new(Arc::new(CapturingAgentFactory { forwarded: forwarded_tx, }))); @@ -1275,11 +1300,16 @@ mod tests { assert_eq!(response.status(), StatusCode::OK); timeout(Duration::from_secs(1), async { - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { connection .push_connection_stream_for_test(format!("message-{index}")) .unwrap(); } + assert!( + connection + .push_connection_stream_for_test("overflow".into()) + .is_err() + ); }) .await .expect("enqueueing must not wait for the SSE body to be polled"); @@ -1297,7 +1327,7 @@ mod tests { .lines() .filter_map(|line| line.strip_prefix("data: ")) .collect::>(); - let expected = (0..ISSUE_288_BURST) + let expected = (0..capacity) .map(|index| format!("message-{index}")) .collect::>(); assert_eq!( @@ -1313,6 +1343,8 @@ mod tests { forwarded: forwarded_tx, }))); let (connection_id, connection) = registry.create_connection().await; + connection.ensure_session("session-1").await; + connection.ensure_session("session-2").await; let request = |session_id: Option<&str>| { let mut request = Request::builder() .method("GET") diff --git a/src/agent-client-protocol-http/src/protocol.rs b/src/agent-client-protocol-http/src/protocol.rs index 79407adf..ff1cc495 100644 --- a/src/agent-client-protocol-http/src/protocol.rs +++ b/src/agent-client-protocol-http/src/protocol.rs @@ -1,4 +1,4 @@ -use agent_client_protocol::{RawJsonRpcMessage, RawJsonRpcParams}; +use agent_client_protocol::{RawJsonRpcMessage, RawJsonRpcParams, schema::v1::RequestId}; pub(crate) const HEADER_CONNECTION_ID: &str = "acp-connection-id"; pub(crate) const HEADER_SESSION_ID: &str = "acp-session-id"; @@ -42,6 +42,19 @@ pub(crate) fn method_for_message(msg: &RawJsonRpcMessage) -> Option<&str> { } } +pub(crate) fn cancelled_request_id(msg: &RawJsonRpcMessage) -> Option { + let RawJsonRpcMessage::Notification(notification) = msg else { + return None; + }; + if notification.method.as_ref() != "$/cancel_request" { + return None; + } + let Some(RawJsonRpcParams::Object(params)) = notification.params.as_ref() else { + return None; + }; + serde_json::from_value(params.get("requestId")?.clone()).ok() +} + pub(crate) fn is_connection_scoped_protocol_message(msg: &RawJsonRpcMessage) -> bool { method_for_message(msg).is_some_and(|method| method.starts_with("$/")) || is_cancel_request_message(msg) diff --git a/src/agent-client-protocol-http/src/websocket_server.rs b/src/agent-client-protocol-http/src/websocket_server.rs index d051f4d0..83882b25 100644 --- a/src/agent-client-protocol-http/src/websocket_server.rs +++ b/src/agent-client-protocol-http/src/websocket_server.rs @@ -235,7 +235,7 @@ where #[cfg(test)] mod tests { use agent_client_protocol::{ - Channel, TransportBatch, TransportBatchEntry, TransportFrame, + BudgetedFrame, Channel, TransportBatch, TransportBatchEntry, TransportFrame, schema::v1::{RequestId, Response as RpcResponse}, }; use async_tungstenite::{tokio::connect_async, tungstenite::Message as ClientWsMessage}; @@ -252,8 +252,6 @@ mod tests { use super::*; - const ISSUE_288_BURST: usize = 1_025; - struct CapturingAgentFactory { forwarded: mpsc::UnboundedSender, } @@ -273,7 +271,7 @@ mod tests { tx: outgoing, } = agent; while let Some(frame) = incoming.next().await { - match frame { + match frame.into_frame() { TransportFrame::Single(message) => { if forwarded.send(message).is_err() { break; @@ -281,9 +279,11 @@ mod tests { } TransportFrame::Malformed { error, .. } => { outgoing - .unbounded_send(TransportFrame::Single( - RawJsonRpcMessage::response(RequestId::Null, Err(error)), - )) + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( + RequestId::Null, + Err(error), + ))) + .await .unwrap(); } TransportFrame::Batch(_) => panic!("expected a single JSON-RPC frame"), @@ -310,7 +310,9 @@ mod tests { let (mut agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); let future = Box::pin(async move { - let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { + let Some(TransportFrame::Batch(batch)) = + agent.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected one batch frame"); }; let mut methods = Vec::new(); @@ -331,7 +333,8 @@ mod tests { TransportBatch::from_messages(responses).expect("responses are non-empty"); agent .tx - .unbounded_send(TransportFrame::Batch(responses)) + .send_frame(TransportFrame::Batch(responses)) + .await .unwrap(); std::future::pending::>().await }); @@ -357,13 +360,14 @@ mod tests { emit.notified().await; agent .tx - .unbounded_send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::notification( "test/final".to_string(), serde_json::json!({}), ) .unwrap(), )) + .await .unwrap(); Ok(()) }); @@ -390,13 +394,14 @@ mod tests { emit.notified().await; agent .tx - .unbounded_send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::notification( "test/final".to_string(), serde_json::json!({}), ) .unwrap(), )) + .await .unwrap(); Ok(()) }); @@ -406,8 +411,9 @@ mod tests { } #[tokio::test] - async fn websocket_buffers_burst_without_polling_slow_subscriber() { + async fn websocket_bounds_burst_and_drains_every_accepted_message() { let (forwarded_tx, _forwarded_rx) = mpsc::unbounded_channel(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; let registry = Arc::new(ConnectionRegistry::new(Arc::new(CapturingAgentFactory { forwarded: forwarded_tx, }))); @@ -425,11 +431,16 @@ mod tests { .await; let mut outbound_rx = connection.subscribe_all_outbound().unwrap(); - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { connection .push_all_outbound_for_test(format!("message-{index}")) .unwrap(); } + assert!( + connection + .push_all_outbound_for_test("overflow".into()) + .is_err() + ); let mut closed = connection.subscribe_closed(); let (mut ws_tx, mut ws_rx) = socket.split(); @@ -457,7 +468,7 @@ mod tests { let (mut client, _) = connect_async(format!("ws://{addr}/acp")).await.unwrap(); timeout(Duration::from_secs(5), async { - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { let frame = client.next().await.unwrap().unwrap(); let ClientWsMessage::Text(text) = frame else { panic!("expected text frame: {frame:?}"); @@ -466,7 +477,7 @@ mod tests { } }) .await - .expect("WebSocket should deliver the complete burst"); + .expect("WebSocket should deliver every accepted frame"); server.abort(); } diff --git a/src/agent-client-protocol-polyfill/Cargo.toml b/src/agent-client-protocol-polyfill/Cargo.toml index 0cb0d67c..cba60257 100644 --- a/src/agent-client-protocol-polyfill/Cargo.toml +++ b/src/agent-client-protocol-polyfill/Cargo.toml @@ -21,7 +21,9 @@ async-stream.workspace = true axum.workspace = true base64.workspace = true futures.workspace = true +hmac = "0.12" serde_json.workspace = true +sha2 = "0.10" tokio = { workspace = true, features = ["net"] } tracing.workspace = true uuid.workspace = true diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs index 9389a81a..3b399685 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs @@ -8,52 +8,85 @@ use std::{convert::Infallible, sync::Arc}; use agent_client_protocol::Error; use axum::{ Json, Router, - body::Bytes, - extract::State, + body::{Body, HttpBody as _, to_bytes}, + extract::{Path, State}, http::{HeaderMap, StatusCode, header}, response::{ IntoResponse, Response, Sse, sse::{Event, KeepAlive}, }, - routing::post, + routing::any, }; use base64::Engine as _; -use futures::{SinkExt, channel::mpsc}; +use futures::{SinkExt, StreamExt, channel::mpsc}; +use hmac::{Hmac, Mac}; use serde_json::{Map, Value}; +use sha2::Sha256; use tokio::{ net::TcpListener, - sync::{mpsc as tokio_mpsc, oneshot}, + sync::{Semaphore, mpsc as tokio_mpsc, oneshot}, }; use super::BridgeMessage; const VERSION: &str = "2026-07-28"; +const MAX_REQUEST_BODY_BYTES: usize = 1024 * 1024; -struct BridgeState { - server_id: String, - token: String, +fn server_route(server_id: &str) -> String { + // Even an empty opaque ID must occupy a real route segment. + format!( + "mcp-{}", + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(server_id) + ) +} + +pub(super) struct BridgeState { + secret: [u8; 32], + admission: Arc, tx: mpsc::Sender, } pub(super) async fn run_http_listener( listener: TcpListener, - server_id: String, - token: String, - tx: mpsc::Sender, + state: Arc, ) -> Result<(), Error> { - let state = Arc::new(BridgeState { - server_id, - token, - tx, - }); let app = Router::new() - .route("/", post(handle_post)) + .route("/{route}", any(handle_request)) .with_state(state); axum::serve(listener, app) .await .map_err(Error::into_internal_error) } +impl BridgeState { + pub(super) fn new(tx: mpsc::Sender) -> Arc { + Arc::new(Self { + secret: { + let mut secret = [0; 32]; + secret[..16].copy_from_slice(uuid::Uuid::new_v4().as_bytes()); + secret[16..].copy_from_slice(uuid::Uuid::new_v4().as_bytes()); + secret + }, + admission: Arc::new(Semaphore::new(super::MAX_ACTIVE_REQUESTS)), + tx, + }) + } + + fn mac(&self, server_id: &str) -> Hmac { + let mut mac = Hmac::::new_from_slice(&self.secret).expect("SHA-256 HMAC key"); + mac.update(b"mcp-over-acp-http-adapter/server/v1\0"); + mac.update(server_id.as_bytes()); + mac + } + + pub(super) fn declaration_url(&self, port: u16, server_id: &str) -> (String, String) { + let route = server_route(server_id); + let token = base64::engine::general_purpose::URL_SAFE_NO_PAD + .encode(self.mac(server_id).finalize().into_bytes()); + (format!("http://127.0.0.1:{port}/{route}"), token) + } +} + fn error(status: StatusCode, id: Value, code: i64, message: &str) -> Response { (status, Json(rpc_error(id, code, message))).into_response() } @@ -89,7 +122,21 @@ pub(super) fn rpc_result(id: Value, request_id: &str, mut result: Value) -> Valu serde_json::json!({"jsonrpc":"2.0", "id":id, "result":result}) } -pub(super) fn rpc_acp_error(id: Value, error: Error) -> Value { +pub(super) fn rpc_binding_error(id: Value, error: Error) -> Value { + let value = serde_json::to_value(error).unwrap_or(Value::Null); + let peer_code = value.get("code").and_then(Value::as_i64); + let code = match peer_code { + Some(-33000 | -33001 | -33002 | -32800) => peer_code.unwrap(), + _ => -33002, + }; + let message = value + .get("message") + .and_then(Value::as_str) + .unwrap_or("MCP binding failure"); + rpc_error(id, code, message) +} + +pub(super) fn rpc_peer_error(id: Value, error: Value) -> Value { serde_json::json!({"jsonrpc":"2.0", "id":id, "error":error}) } @@ -114,14 +161,31 @@ fn valid_origin(headers: &HeaderMap) -> bool { } fn accepts_both(headers: &HeaderMap) -> bool { - let Some(accept) = header_value(headers, "accept") else { - return false; - }; - let types = accept - .split(',') - .map(|part| part.split(';').next().unwrap_or("").trim()); - let types: Vec<_> = types.collect(); - types.contains(&"application/json") && types.contains(&"text/event-stream") + let mut json = false; + let mut sse = false; + for value in headers.get_all(header::ACCEPT) { + let Ok(value) = value.to_str() else { + return false; + }; + for item in value.split(',') { + let mut parts = item.split(';'); + let media = parts.next().unwrap_or("").trim(); + let mut quality = 1.0; + for part in parts { + if let Some((key, q)) = part.trim().split_once('=') + && key.trim().eq_ignore_ascii_case("q") + { + quality = q.trim().parse::().unwrap_or(0.0); + } + } + if quality <= 0.0 || quality > 1.0 { + continue; + } + json |= media.eq_ignore_ascii_case("application/json"); + sse |= media.eq_ignore_ascii_case("text/event-stream"); + } + } + json && sse } fn mirrored_name<'a>(method: &str, params: &'a Map) -> Option<&'a str> { @@ -147,14 +211,16 @@ fn matches_mirror(header: Option<&str>, body: &str) -> bool { .is_ok_and(|bytes| bytes == body.as_bytes()) } else { // Literal sentinel-looking values must be encoded to avoid ambiguity. - !header.starts_with("=?base64?") && header == body + !(header.starts_with("=?base64?") && header.ends_with("?=")) && header == body } } -async fn handle_post( +async fn handle_request( State(state): State>, + Path(route): Path, + method: axum::http::Method, headers: HeaderMap, - body: Bytes, + body: Body, ) -> Response { if [ "origin", @@ -176,12 +242,59 @@ async fn handle_post( if !valid_origin(&headers) { return error(StatusCode::FORBIDDEN, Value::Null, -32600, "Invalid Origin"); } - if header_value(&headers, "authorization") != Some(&format!("Bearer {}", state.token)) { + if route.len() > 4096 { + return error( + StatusCode::NOT_FOUND, + Value::Null, + -32601, + "Unknown MCP route", + ); + } + let server_id = { + let decoded = route.strip_prefix("mcp-").and_then(|encoded| { + base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(encoded) + .ok() + }); + let Some(server_id) = decoded + .and_then(|id| String::from_utf8(id).ok()) + .filter(|id| server_route(id) == route) + else { + return error( + StatusCode::NOT_FOUND, + Value::Null, + -32601, + "Unknown MCP route", + ); + }; + let authorization = + header_value(&headers, "authorization").and_then(|value| value.split_once(' ')); + if !authorization.is_some_and(|(scheme, supplied)| { + scheme.eq_ignore_ascii_case("bearer") + && base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(supplied) + .is_ok_and(|tag| state.mac(&server_id).verify_slice(&tag).is_ok()) + }) { + let mut response = error( + StatusCode::UNAUTHORIZED, + Value::Null, + -32600, + "Unauthorized", + ); + response.headers_mut().insert( + header::WWW_AUTHENTICATE, + "Bearer".parse().expect("static header"), + ); + return response; + } + server_id + }; + if method != axum::http::Method::POST { return error( - StatusCode::UNAUTHORIZED, + StatusCode::METHOD_NOT_ALLOWED, Value::Null, -32600, - "Unauthorized", + "Only POST is supported", ); } if !accepts_both(&headers) { @@ -192,9 +305,12 @@ async fn handle_post( "Accept must include application/json and text/event-stream", ); } - if header_value(&headers, header::CONTENT_TYPE.as_str()) - .is_none_or(|value| !value.eq_ignore_ascii_case("application/json")) - { + if header_value(&headers, header::CONTENT_TYPE.as_str()).is_none_or(|value| { + !value + .split(';') + .next() + .is_some_and(|media| media.trim().eq_ignore_ascii_case("application/json")) + }) { return error( StatusCode::UNSUPPORTED_MEDIA_TYPE, Value::Null, @@ -202,6 +318,50 @@ async fn handle_post( "Expected application/json", ); } + // Acquire before reading a potentially slow/large request body. The permit + // stays owned by the response body until the client consumes or drops it. + let Ok(permit) = state.admission.clone().try_acquire_owned() else { + return error( + StatusCode::TOO_MANY_REQUESTS, + Value::Null, + -33000, + "Too many outstanding MCP responses", + ); + }; + let response = handle_admitted_request(state, server_id, headers, body).await; + let (mut parts, body) = response.into_parts(); + if let Some(length) = body.size_hint().exact() { + parts + .headers + .entry(header::CONTENT_LENGTH) + .or_insert_with(|| length.to_string().parse().expect("decimal body length")); + } + // One ownership rule for every admitted response, including validation + // failures that echo a potentially large, but valid, external request ID. + let stream = async_stream::stream! { + let _permit = permit; + let mut body = body.into_data_stream(); + while let Some(chunk) = body.next().await { + yield chunk; + } + }; + Response::from_parts(parts, Body::from_stream(stream)) +} + +async fn handle_admitted_request( + state: Arc, + server_id: String, + headers: HeaderMap, + body: Body, +) -> Response { + let Ok(body) = to_bytes(body, MAX_REQUEST_BODY_BYTES).await else { + return error( + StatusCode::PAYLOAD_TOO_LARGE, + Value::Null, + -33000, + "Request body too large", + ); + }; let body: Value = match serde_json::from_slice(&body) { Ok(body) => body, Err(_) => return error(StatusCode::BAD_REQUEST, Value::Null, -32700, "Parse error"), @@ -301,9 +461,8 @@ async fn handle_post( ); } } - // Tool schemas with x-mcp-header annotations are not tracked in this adapter. - // Fail closed on supplied mirrored parameter headers; support for annotations - // requires a request-scoped schema lookup and validation before forwarding. + // This endpoint re-exports native tools without transport-only x-mcp-header + // annotations. Mirrored parameter headers have no authority here. if headers .keys() .any(|key| key.as_str().starts_with("mcp-param-")) @@ -318,6 +477,7 @@ async fn handle_post( if method == "initialize" || method.starts_with("notifications/") { return error(StatusCode::NOT_FOUND, id, -32601, "Method not found"); } + let id_for_bridge_error = id.clone(); let (notification_tx, mut response_rx) = tokio_mpsc::channel(super::MAX_QUEUED_NOTIFICATIONS); let response_tx = super::StreamSender { tx: notification_tx, @@ -325,7 +485,7 @@ async fn handle_post( }; let (terminal_tx, mut terminal_rx) = oneshot::channel(); let message = BridgeMessage::Request { - server_id: state.server_id.clone(), + server_id, request_id: uuid::Uuid::new_v4().to_string(), http_id: id, method: method.into(), @@ -337,8 +497,8 @@ async fn handle_post( if tx.send(message).await.is_err() { return error( StatusCode::SERVICE_UNAVAILABLE, - Value::Null, - -32603, + id_for_bridge_error.clone(), + -33002, "ACP bridge unavailable", ); } @@ -354,8 +514,8 @@ async fn handle_post( let Some(first) = first else { return error( StatusCode::SERVICE_UNAVAILABLE, - Value::Null, - -32603, + id_for_bridge_error, + -33002, "ACP bridge closed", ); }; @@ -365,7 +525,20 @@ async fn handle_post( } else { StatusCode::OK }; - return (status, Json(first)).into_response(); + let payload = first.to_string(); + let length = payload.len().to_string(); + let stream = async_stream::stream! { + yield Ok::<_, Infallible>(axum::body::Bytes::from(payload)); + }; + return ( + status, + [ + (header::CONTENT_TYPE, "application/json".to_string()), + (header::CONTENT_LENGTH, length), + ], + Body::from_stream(stream), + ) + .into_response(); } let stream = async_stream::stream! { yield Ok::<_, Infallible>(Event::default().data(first.to_string())); @@ -401,6 +574,169 @@ mod tests { use super::*; use tokio::io::{AsyncReadExt, AsyncWriteExt}; + #[test] + fn stateless_declarations_do_not_allocate_routes() { + let (tx, _rx) = mpsc::channel(1); + let state = BridgeState::new(tx); + let first = state.declaration_url(1234, "server/one"); + let other = state.declaration_url(1234, "server/two"); + assert_eq!(first, state.declaration_url(1234, "server/one")); + assert_ne!(first, other); + assert!(state.declaration_url(1234, "").0.ends_with("/mcp-")); + for i in 0..1000 { + let (url, bearer) = state.declaration_url(1234, &i.to_string()); + assert!(url.starts_with("http://127.0.0.1:1234/")); + assert!(!url.contains(&bearer)); + } + } + + #[tokio::test] + async fn bridge_failure_preserves_valid_external_id() { + let (tx, rx) = mpsc::channel(1); + let state = BridgeState::new(tx); + drop(rx); + let (_, token) = state.declaration_url(8000, "server"); + let mut headers = HeaderMap::new(); + headers.insert("host", "127.0.0.1:8000".parse().unwrap()); + headers.insert("authorization", format!("Bearer {token}").parse().unwrap()); + headers.insert( + "accept", + "application/json, text/event-stream".parse().unwrap(), + ); + headers.insert("content-type", "application/json".parse().unwrap()); + headers.insert("mcp-protocol-version", VERSION.parse().unwrap()); + headers.insert("mcp-method", "tools/list".parse().unwrap()); + let response = handle_request( + State(state), + Path(server_route("server")), + axum::http::Method::POST, + headers, + Body::from( + serde_json::json!({"jsonrpc":"2.0","id":"external", + "method":"tools/list","params":{"_meta":{ + "io.modelcontextprotocol/protocolVersion":VERSION, + "io.modelcontextprotocol/clientCapabilities":{}}}}) + .to_string(), + ), + ) + .await; + assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + let bytes = to_bytes(response.into_body(), MAX_REQUEST_BODY_BYTES) + .await + .unwrap(); + let body: Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(body["id"], "external"); + assert_eq!(body["error"]["code"], -33002); + } + + #[tokio::test] + async fn unread_validation_errors_hold_admission_until_consumed_or_dropped() { + let (tx, _rx) = mpsc::channel(1); + let state = BridgeState::new(tx); + let (_, token) = state.declaration_url(8000, "server"); + let route = server_route("server"); + let mut headers = HeaderMap::new(); + headers.insert("authorization", format!("Bearer {token}").parse().unwrap()); + headers.insert( + "accept", + "application/json, text/event-stream".parse().unwrap(), + ); + headers.insert("content-type", "application/json".parse().unwrap()); + // No method: validation must echo this large known ID without releasing + // the permit while the client still owns its unread response. + let id = "external".repeat(32 * 1024); + let body = serde_json::json!({"jsonrpc":"2.0", "id":id}).to_string(); + let send = || { + handle_request( + State(state.clone()), + Path(route.clone()), + axum::http::Method::POST, + headers.clone(), + Body::from(body.clone()), + ) + }; + let mut responses = Vec::new(); + for _ in 0..super::super::MAX_ACTIVE_REQUESTS { + let response = send().await; + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + responses.push(response); + } + assert_eq!(send().await.status(), StatusCode::TOO_MANY_REQUESTS); + + let bytes = to_bytes(responses.pop().unwrap().into_body(), MAX_REQUEST_BODY_BYTES) + .await + .unwrap(); + let error: Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(error["id"], id); + assert_eq!(error["error"]["code"], -32600); + assert_eq!(state.admission.available_permits(), 1); + responses.push(send().await); + assert_eq!(state.admission.available_permits(), 0); + + drop(responses.pop()); + assert_eq!(state.admission.available_permits(), 1); + let recovered = send().await; + assert_eq!(recovered.status(), StatusCode::BAD_REQUEST); + drop(recovered); + drop(responses); + assert_eq!( + state.admission.available_permits(), + super::super::MAX_ACTIVE_REQUESTS + ); + } + + #[tokio::test] + async fn unread_terminal_bodies_hold_admission_until_drop() { + let (tx, mut rx) = mpsc::channel(128); + let state = BridgeState::new(tx); + let (_, token) = state.declaration_url(8000, "server"); + let route = server_route("server"); + let mut headers = HeaderMap::new(); + headers.insert("host", "127.0.0.1:8000".parse().unwrap()); + headers.insert("authorization", format!("bearer {token}").parse().unwrap()); + headers.insert("accept", "application/json".parse().unwrap()); + headers.append("accept", "text/event-stream;q=0.8".parse().unwrap()); + headers.insert( + "content-type", + "application/json; charset=utf-8".parse().unwrap(), + ); + headers.insert("mcp-protocol-version", VERSION.parse().unwrap()); + headers.insert("mcp-method", "tools/list".parse().unwrap()); + tokio::spawn(async move { + while let Some(BridgeMessage::Request { + terminal_tx, + http_id, + .. + }) = rx.next().await + { + drop(terminal_tx.send(rpc_result(http_id, "", serde_json::json!({"tools":[]})))); + } + }); + let body = serde_json::json!({"jsonrpc":"2.0","id":"known","method":"tools/list", + "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":VERSION, + "io.modelcontextprotocol/clientCapabilities":{}}}}) + .to_string(); + let send = || { + handle_request( + State(state.clone()), + Path(route.clone()), + axum::http::Method::POST, + headers.clone(), + Body::from(body.clone()), + ) + }; + let mut responses = Vec::new(); + for _ in 0..super::super::MAX_ACTIVE_REQUESTS { + let response = send().await; + assert_eq!(response.status(), StatusCode::OK); + responses.push(response); + } + assert_eq!(send().await.status(), StatusCode::TOO_MANY_REQUESTS); + drop(responses.pop()); + let recovered = send().await; + assert_eq!(recovered.status(), StatusCode::OK); + } + #[test] fn accepts_only_both_media_types() { let mut headers = HeaderMap::new(); @@ -411,6 +747,13 @@ mod tests { assert!(accepts_both(&headers)); headers.insert("accept", "application/json".parse().unwrap()); assert!(!accepts_both(&headers)); + headers.append("accept", "text/event-stream;q=0.9".parse().unwrap()); + assert!(accepts_both(&headers)); + headers.insert( + "accept", + "application/json, text/event-stream;q=0".parse().unwrap(), + ); + assert!(!accepts_both(&headers)); } #[test] @@ -426,6 +769,10 @@ mod tests { Some("=?base64?literal?="), "=?base64?literal?=" )); + assert!(matches_mirror( + Some("=?base64?unfinished"), + "=?base64?unfinished" + )); } #[test] @@ -478,13 +825,14 @@ mod tests { async fn rejects_legacy_methods_and_invalid_headers_over_real_http() { async fn exchange( address: std::net::SocketAddr, + route: &str, method: &str, headers: &str, body: &str, ) -> String { let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); let request = format!( - "{method} / HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n{headers}Content-Length: {}\r\n\r\n{body}", + "{method} /{route} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n{headers}Content-Length: {}\r\n\r\n{body}", body.len() ); stream.write_all(request.as_bytes()).await.unwrap(); @@ -495,38 +843,52 @@ mod tests { let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let address = listener.local_addr().unwrap(); let (tx, _rx) = mpsc::channel(8); - let task = tokio::spawn(run_http_listener( - listener, - "server".into(), - "secret".into(), - tx, - )); - let legacy = exchange(address, "GET", "", "").await; + let state = BridgeState::new(tx); + let (url, token) = state.declaration_url(address.port(), "server"); + let route = url.rsplit('/').next().unwrap(); + let task = tokio::spawn(run_http_listener(listener, state)); + let auth = format!("Authorization: Bearer {token}\r\n"); + let legacy = exchange(address, route, "GET", &auth, "").await; assert!(legacy.starts_with("HTTP/1.1 405"), "{legacy}"); - let delete = exchange(address, "DELETE", "", "").await; + let delete = exchange(address, route, "DELETE", &auth, "").await; assert!(delete.starts_with("HTTP/1.1 405"), "{delete}"); - let invalid_origin = exchange(address, "POST", "Origin: http://evil.test\r\n", "{}").await; + let invalid_origin = + exchange(address, route, "POST", "Origin: http://evil.test\r\n", "{}").await; assert!( invalid_origin.starts_with("HTTP/1.1 403"), "{invalid_origin}" ); - let invalid_auth = exchange(address, "POST", "", "{}").await; + let invalid_get_origin = + exchange(address, route, "GET", "Origin: http://evil.test\r\n", "").await; + assert!( + invalid_get_origin.starts_with("HTTP/1.1 403"), + "{invalid_get_origin}" + ); + let invalid_auth = exchange(address, route, "POST", "", "{}").await; assert!(invalid_auth.starts_with("HTTP/1.1 401"), "{invalid_auth}"); + assert!( + invalid_auth + .to_ascii_lowercase() + .contains("www-authenticate: bearer"), + "{invalid_auth}" + ); let body = serde_json::json!({"jsonrpc":"2.0","id":1,"method":"tools/list", "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":VERSION, "io.modelcontextprotocol/clientCapabilities":{}}}}) .to_string(); - let headers = "Authorization: Bearer secret\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: wrong/method\r\n"; - let mismatch = exchange(address, "POST", headers, &body).await; + let headers = format!( + "{auth}Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: wrong/method\r\n" + ); + let mismatch = exchange(address, route, "POST", &headers, &body).await; assert!(mismatch.starts_with("HTTP/1.1 400"), "{mismatch}"); assert!(mismatch.contains("-32020"), "{mismatch}"); - let batch = exchange(address, "POST", headers, "[]").await; + let batch = exchange(address, route, "POST", &headers, "[]").await; assert!(batch.starts_with("HTTP/1.1 400"), "{batch}"); let headers = headers.replace("wrong/method", "tools/list"); let fractional_id = body.replace("\"id\":1", "\"id\":1.5"); let fractional = tokio::time::timeout( std::time::Duration::from_secs(3), - exchange(address, "POST", &headers, &fractional_id), + exchange(address, route, "POST", &headers, &fractional_id), ) .await .expect("an invalid request ID must be rejected before forwarding"); diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs index 040f4f6d..0b0262ee 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs @@ -7,7 +7,7 @@ pub(crate) mod http; mod protocol; use std::{ - collections::{HashMap, HashSet}, + collections::HashMap, sync::{ Arc, atomic::{AtomicUsize, Ordering}, @@ -16,7 +16,7 @@ use std::{ use agent_client_protocol::{ Agent, Client, Conductor, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - Proxy, UntypedMessage, util::MatchDispatchFrom, + Proxy, UntypedMessage, schema::v1::MessageMcpResponse, util::MatchDispatchFrom, }; use futures::{ SinkExt, StreamExt, @@ -26,14 +26,15 @@ use serde_json::Value; use tokio::{net::TcpListener, sync::mpsc as tokio_mpsc}; use tracing::{debug, warn}; -use self::protocol::{DownstreamMcpMode, NativeMcpNotification, NativeServer, PolyfillProtocol}; +use self::protocol::{DownstreamMcpMode, NativeMcpNotification, PolyfillProtocol}; // Conservative per-bridge limits. Notifications are bounded per HTTP POST by // both message count and serialized bytes; terminal responses bypass the queue. const MAX_ACTIVE_REQUESTS: usize = 64; -const MAX_LISTENERS: usize = 32; const MAX_QUEUED_NOTIFICATIONS: usize = 16; const MAX_QUEUED_BYTES: usize = 256 * 1024; +const MAX_TERMINAL_BYTES: usize = 1024 * 1024; +const LOCAL_LIMIT_ERROR: i64 = -33000; struct QueuedNotification { value: Value, @@ -152,7 +153,7 @@ impl ConnectTo for McpOverAcpProxy { bridge_rx, protocol: None, downstream_mode: DownstreamMcpMode::Unknown, - listeners: HashMap::new(), + listener: None, active: HashMap::new(), }; let handler = PolyfillHandler { @@ -325,26 +326,6 @@ async fn transform_session_servers( Ok(()) } -struct BridgeListener { - tcp_port: u16, - // Runtime-only; never trace the listener or the rewritten declaration. - token: String, -} - -impl BridgeListener { - fn declaration( - &self, - protocol: PolyfillProtocol, - server: NativeServer, - ) -> Result { - server.http_declaration( - protocol, - format!("http://127.0.0.1:{}", self.tcp_port), - &self.token, - ) - } -} - struct ActiveRequest { server_id: String, http_id: Value, @@ -359,7 +340,7 @@ struct BridgeRunner { bridge_rx: mpsc::Receiver, protocol: Option, downstream_mode: DownstreamMcpMode, - listeners: HashMap, + listener: Option<(u16, Arc)>, active: HashMap, } @@ -368,7 +349,7 @@ impl std::fmt::Debug for BridgeRunner { f.debug_struct("BridgeRunner") .field("protocol", &self.protocol) .field("downstream_mode", &self.downstream_mode) - .field("listeners", &self.listeners.len()) + .field("listener", &self.listener.is_some()) .field("active", &self.active.len()) .finish_non_exhaustive() } @@ -410,23 +391,15 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { else { drop(terminal_tx.send(http::rpc_error( http_id, - -32603, + -33002, "MCP adapter unavailable", ))); continue; }; - if !self.listeners.contains_key(&server_id) { - drop(terminal_tx.send(http::rpc_error( - http_id, - -32602, - "Unknown MCP server", - ))); - continue; - } if !self.can_admit_request() { drop(terminal_tx.send(http::rpc_error( http_id, - -32000, + LOCAL_LIMIT_ERROR, "Too many active MCP requests", ))); continue; @@ -493,7 +466,7 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { let _ = active.cancel_tx.send(()); drop(active.terminal_tx.send(http::rpc_error( active.http_id, - -32000, + LOCAL_LIMIT_ERROR, "MCP notification queue overflow", ))); } @@ -504,14 +477,26 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { continue; }; if let Some(result) = result { + let http_id = active.http_id.clone(); let value = match result { - Ok(mut result) => { - if active.method == "tools/list" { - filter_annotated_tools(&mut result); - } - http::rpc_result(active.http_id, &request_id, result) - } - Err(error) => http::rpc_acp_error(active.http_id, error), + Ok(carrier) => project_mcp_carrier( + active.http_id, + &request_id, + &active.method, + carrier, + ), + Err(error) => http::rpc_binding_error(active.http_id, error), + }; + let value = if serde_json::to_vec(&value) + .is_ok_and(|bytes| bytes.len() <= MAX_TERMINAL_BYTES) + { + value + } else { + http::rpc_error( + http_id, + LOCAL_LIMIT_ERROR, + "MCP terminal response too large", + ) }; drop(active.terminal_tx.send(value)); } @@ -522,6 +507,26 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { } } +/// ACP success carries exactly one MCP outcome. An outer ACP failure is a +/// binding/runtime failure, not an MCP error carried in a successful response. +fn project_mcp_carrier(http_id: Value, request_id: &str, method: &str, carrier: Value) -> Value { + // Both ACP revisions share this type. Keep envelope validation in the schema, + // rather than maintaining a second parser that can drift from its null rules. + match serde_json::from_value::(carrier) { + Ok(MessageMcpResponse::Result { mut result, .. }) => { + if method == "tools/list" { + strip_header_annotations(&mut result); + } + http::rpc_result(http_id, request_id, result) + } + Ok(MessageMcpResponse::Error { error, .. }) => http::rpc_peer_error( + http_id, + serde_json::to_value(error).expect("MCP errors contain only JSON values"), + ), + _ => http::rpc_error(http_id, -33002, "Invalid MCP-over-ACP response carrier"), + } +} + impl BridgeRunner { fn can_admit_request(&self) -> bool { self.active.len() < MAX_ACTIVE_REQUESTS @@ -544,11 +549,7 @@ impl BridgeRunner { match self.downstream_mode { DownstreamMcpMode::Native => transformed.push(server), DownstreamMcpMode::HttpAdapter => { - if !self.listeners.contains_key(&native.server_id) { - if self.listeners.len() >= MAX_LISTENERS { - return Err(agent_client_protocol::Error::invalid_params() - .data("too many MCP HTTP listeners")); - } + if self.listener.is_none() { let listener = TcpListener::bind("127.0.0.1:0") .await .map_err(agent_client_protocol::Error::into_internal_error)?; @@ -556,28 +557,13 @@ impl BridgeRunner { .local_addr() .map_err(agent_client_protocol::Error::into_internal_error)? .port(); - let token = uuid::Uuid::new_v4().simple().to_string() - + &uuid::Uuid::new_v4().simple().to_string(); - connection.spawn(http::run_http_listener( - listener, - native.server_id.clone(), - token.clone(), - self.bridge_tx.clone(), - ))?; - self.listeners.insert( - native.server_id.clone(), - BridgeListener { - tcp_port: port, - token, - }, - ); + let state = http::BridgeState::new(self.bridge_tx.clone()); + connection.spawn(http::run_http_listener(listener, state.clone()))?; + self.listener = Some((port, state)); } - transformed.push( - self.listeners - .get(&native.server_id) - .expect("listener created") - .declaration(protocol, native)?, - ); + let (port, state) = self.listener.as_ref().expect("listener created"); + let (url, token) = state.declaration_url(*port, &native.server_id); + transformed.push(native.http_declaration(protocol, url, &token)?); } DownstreamMcpMode::Unknown | DownstreamMcpMode::Unavailable => { return Err(agent_client_protocol::Error::invalid_params().data( @@ -590,9 +576,6 @@ impl BridgeRunner { } } -/// For each tools/call POST, inspect the current tool schema in that request's -/// scope. This adds an ACP tools/list lookup, but requires no client-side -/// discovery handshake and cannot silently omit an annotated parameter header. async fn forward_http_request( connection: ConnectionTo, protocol: PolyfillProtocol, @@ -601,65 +584,6 @@ async fn forward_http_request( method: String, params: Option>, ) -> Result { - if method == "tools/call" { - let name = params - .as_ref() - .and_then(|p| p.get("name")) - .and_then(Value::as_str) - .ok_or_else(agent_client_protocol::Error::invalid_params)?; - let meta = params.as_ref().and_then(|p| p.get("_meta")).cloned(); - let mut cursor: Option = None; - let mut seen = HashSet::new(); - loop { - let mut list_params = serde_json::Map::new(); - if let Some(meta) = &meta { - list_params.insert("_meta".into(), meta.clone()); - } - if let Some(cursor) = &cursor { - list_params.insert("cursor".into(), Value::String(cursor.clone())); - } - let lookup = protocol.message_request( - server_id.clone(), - uuid::Uuid::new_v4().to_string(), - "tools/list".into(), - Some(list_params), - None, - )?; - let listing = connection - .send_request_to(Client, lookup) - .block_task() - .await?; - let tools = listing - .get("tools") - .and_then(Value::as_array) - .ok_or_else(|| { - agent_client_protocol::Error::invalid_params() - .data("tools/list result must contain a tools array") - })?; - if let Some(tool) = tools - .iter() - .find(|tool| tool.get("name").and_then(Value::as_str) == Some(name)) - { - if tool - .get("inputSchema") - .is_none_or(|schema| !schema.is_object() || contains_header_annotation(schema)) - { - return Err(agent_client_protocol::Error::invalid_params() - .data("tool uses x-mcp-header or has no verifiable input schema")); - } - break; - } - let Some(next) = listing.get("nextCursor").and_then(Value::as_str) else { - return Err(agent_client_protocol::Error::invalid_params() - .data("tool was not found in tools/list")); - }; - if !seen.insert(next.to_owned()) || seen.len() > 128 { - return Err(agent_client_protocol::Error::invalid_params() - .data("tools/list pagination did not terminate")); - } - cursor = Some(next.to_owned()); - } - } let request = protocol.message_request(server_id, request_id, method, params, None)?; connection .send_request_to(Client, request) @@ -667,42 +591,95 @@ async fn forward_http_request( .await } -fn contains_header_annotation(value: &Value) -> bool { - match value { - Value::Object(object) => { - object.contains_key("x-mcp-header") || object.values().any(contains_header_annotation) +fn strip_header_annotations(result: &mut Value) { + let Some(tools) = result.get_mut("tools").and_then(Value::as_array_mut) else { + return; + }; + for tool in tools { + if let Some(schema) = tool.get_mut("inputSchema") { + strip_schema_annotation(schema); } - Value::Array(values) => values.iter().any(contains_header_annotation), - _ => false, } } -fn filter_annotated_tools(result: &mut Value) { - let Some(tools) = result.get_mut("tools").and_then(Value::as_array_mut) else { +fn strip_schema_annotation(schema: &mut Value) { + let Some(object) = schema.as_object_mut() else { return; }; - tools.retain(|tool| { - let Some(name) = tool.get("name").and_then(Value::as_str) else { - return false; - }; - if tool - .get("inputSchema") - .is_none_or(|schema| !schema.is_object() || contains_header_annotation(schema)) - { - warn!( - tool = name, - "excluding tool with unsupported x-mcp-header annotation" - ); - return false; + object.remove("x-mcp-header"); + for key in [ + "properties", + "patternProperties", + "$defs", + "definitions", + "dependentSchemas", + ] { + if let Some(children) = object.get_mut(key).and_then(Value::as_object_mut) { + for child in children.values_mut() { + strip_schema_annotation(child); + } + } + } + for key in [ + "items", + "additionalItems", + "additionalProperties", + "unevaluatedItems", + "unevaluatedProperties", + "contains", + "contentSchema", + "not", + "if", + "then", + "else", + "propertyNames", + ] { + if let Some(child) = object.get_mut(key) { + strip_schema_annotation(child); + } + } + for key in ["allOf", "anyOf", "oneOf", "prefixItems"] { + if let Some(children) = object.get_mut(key).and_then(Value::as_array_mut) { + for child in children { + strip_schema_annotation(child); + } } - true - }); + } } #[cfg(test)] mod http_limits_tests { use super::*; + #[test] + fn annotation_removal_only_traverses_schema_locations() { + let mut listing = serde_json::json!({"tools":[{ + "name":"with-header", + "inputSchema":{ + "type":"object", + "properties":{ + "x-mcp-header":{"type":"string","default":"retain"}, + "nested":{"type":"object","x-mcp-header":"Nested","properties":{ + "value":{"type":"string","x-mcp-header":"Value", + "examples":[{"x-mcp-header":"user data"}]} + }} + }, + "$defs":{"inner":{"type":"string","x-mcp-header":"Inner"}}, + "default":{"x-mcp-header":"not a schema"} + } + }]}); + strip_header_annotations(&mut listing); + let schema = &listing["tools"][0]["inputSchema"]; + assert_eq!(schema["properties"]["x-mcp-header"]["default"], "retain"); + assert_eq!( + schema["properties"]["nested"]["properties"]["value"]["examples"][0]["x-mcp-header"], + "user data" + ); + assert_eq!(schema["default"]["x-mcp-header"], "not a schema"); + assert!(schema["properties"]["nested"].get("x-mcp-header").is_none()); + assert!(schema["$defs"]["inner"].get("x-mcp-header").is_none()); + } + #[test] fn slow_reader_overflows_by_count_without_blocking_other_requests() { let (tx, mut rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); @@ -755,14 +732,14 @@ mod http_limits_tests { } #[test] - fn admission_reopens_when_an_active_request_finishes() { + fn backend_capacity_reopens_when_an_active_request_finishes() { let (bridge_tx, bridge_rx) = mpsc::channel(1); let mut runner = BridgeRunner { bridge_tx, bridge_rx, protocol: None, downstream_mode: DownstreamMcpMode::Unknown, - listeners: HashMap::new(), + listener: None, active: HashMap::new(), }; let (tx, _rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); @@ -796,15 +773,50 @@ mod tests { use super::*; #[test] - fn annotated_tools_are_not_advertised_or_callable() { + fn mcp_carrier_preserves_peer_error_and_rejects_ambiguous_outcomes() { + let error = serde_json::json!({ + "code":-32000,"message":"peer-defined error", + "data":{"nested":[1,2]},"extension":"preserved" + }); + let project = |carrier| { + project_mcp_carrier( + serde_json::json!("external"), + "internal", + "tools/call", + carrier, + ) + }; + assert_eq!(project(serde_json::json!({"error":error}))["error"], error); + assert_eq!( + project(serde_json::json!({"result":null})), + serde_json::json!({"jsonrpc":"2.0","id":"external","result":null}) + ); + for invalid in [ + serde_json::json!({"result":null,"error":error}), + serde_json::json!({"tools":[]}), + serde_json::json!({"error":null}), + ] { + let response = project(invalid); + assert_eq!(response["id"], "external"); + assert_eq!(response["error"]["code"], -33002); + } + } + + #[test] + fn annotated_tools_are_reexported_without_transport_annotations() { let mut result = serde_json::json!({"tools":[ {"name":"plain","inputSchema":{"type":"object","properties":{}}}, {"name":"annotated","inputSchema":{"properties":{"nested":{"properties":{ "region":{"type":"string","x-mcp-header":"Region"} }}}}} ]}); - filter_annotated_tools(&mut result); - assert_eq!(result["tools"].as_array().unwrap().len(), 1); + strip_header_annotations(&mut result); + assert_eq!(result["tools"].as_array().unwrap().len(), 2); assert_eq!(result["tools"][0]["name"], "plain"); + assert_eq!(result["tools"][1]["name"], "annotated"); + assert_eq!( + result["tools"][1]["inputSchema"]["properties"]["nested"]["properties"]["region"], + serde_json::json!({"type":"string"}) + ); } } diff --git a/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs b/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs index 81a206fe..8263459f 100644 --- a/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs +++ b/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs @@ -116,7 +116,13 @@ async fn main() -> Result<(), Error> { .block_task() .await?; let response = done_rx.await.map_err(Error::into_internal_error)??; - println!("{}", response.0.get()); + match response { + v2::MessageMcpResponse::Result { result, .. } => println!("{result}"), + v2::MessageMcpResponse::Error { error, .. } => { + eprintln!("MCP error {}: {}", error.code, error.message); + } + _ => return Err(Error::internal_error().data("unknown MCP carrier")), + } Ok(()) }) .await diff --git a/src/agent-client-protocol-rmcp/src/builder.rs b/src/agent-client-protocol-rmcp/src/builder.rs index 585cd94c..6f8d2af4 100644 --- a/src/agent-client-protocol-rmcp/src/builder.rs +++ b/src/agent-client-protocol-rmcp/src/builder.rs @@ -15,6 +15,8 @@ use schemars::JsonSchema; use serde::{Serialize, de::DeserializeOwned}; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; +#[cfg(feature = "unstable_mcp_over_acp")] +use acp::mcp_server::{McpOutcome, McpRequest, McpRequestContext, McpService}; use agent_client_protocol as acp; use agent_client_protocol::{ ByteStreams, ChainRun, ConnectTo, DynConnectTo, NullRun, RunWithConnectionTo, @@ -231,13 +233,22 @@ where /// feature, it can also be attached through /// `SessionBuilder::with_mcp_server` or `Builder::with_mcp_server`. pub fn build(self) -> McpServer { - McpServer::new( - McpServerBuilt { - name: self.name, - data: Arc::new(self.data), - }, - self.runner, - ) + let built = McpServerBuilt { + name: self.name, + data: Arc::new(self.data), + }; + #[cfg(feature = "unstable_mcp_over_acp")] + { + let standalone = McpServerBuilt { + name: built.name.clone(), + data: built.data.clone(), + }; + McpServer::new_service_with_standalone(built, standalone, self.runner) + } + #[cfg(not(feature = "unstable_mcp_over_acp"))] + { + McpServer::new(built, self.runner) + } } } @@ -246,6 +257,21 @@ struct McpServerBuilt { data: Arc>, } +#[cfg(feature = "unstable_mcp_over_acp")] +impl McpService for McpServerBuilt { + fn execute( + &self, + request: McpRequest, + context: McpRequestContext, + ) -> BoxFuture<'static, Result> { + let handler = McpServerConnection { + data: self.data.clone(), + mcp_connection: context.connection().clone(), + }; + crate::native::execute(Arc::new(handler), request, context) + } +} + impl McpServerConnect for McpServerBuilt { fn name(&self) -> String { self.name.clone() diff --git a/src/agent-client-protocol-rmcp/src/lib.rs b/src/agent-client-protocol-rmcp/src/lib.rs index 1a91a54d..34f1f5e2 100644 --- a/src/agent-client-protocol-rmcp/src/lib.rs +++ b/src/agent-client-protocol-rmcp/src/lib.rs @@ -40,13 +40,21 @@ //! ``` use agent_client_protocol::mcp_server::{McpConnectionTo, McpServer, McpServerConnect}; +#[cfg(feature = "unstable_mcp_over_acp")] +use agent_client_protocol::mcp_server::{McpOutcome, McpRequest, McpRequestContext, McpService}; use agent_client_protocol::role; use agent_client_protocol::{ByteStreams, ConnectTo, DynConnectTo, NullRun, Role}; +#[cfg(feature = "unstable_mcp_over_acp")] +use futures::future::BoxFuture; use futures_concurrency::future::TryJoin as _; use rmcp::ServiceExt; +#[cfg(feature = "unstable_mcp_over_acp")] +use std::sync::Arc; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; mod builder; +#[cfg(feature = "unstable_mcp_over_acp")] +mod native; pub use agent_client_protocol::mcp_server::{EnabledTools, McpTool}; pub use agent_client_protocol::{tool_fn, tool_fn_mut}; @@ -76,6 +84,29 @@ pub trait McpServerExt { new_fn: F, } + #[cfg(feature = "unstable_mcp_over_acp")] + struct SharedRmcp { + new_fn: Arc, + service: std::sync::OnceLock>, + } + + #[cfg(feature = "unstable_mcp_over_acp")] + impl McpService for SharedRmcp + where + Counterpart: Role, + F: Fn() -> S + Send + Sync + 'static, + S: rmcp::Service, + { + fn execute( + &self, + request: McpRequest, + context: McpRequestContext, + ) -> BoxFuture<'static, Result> { + let service = self.service.get_or_init(|| Arc::new((self.new_fn)())); + native::execute(service.clone(), request, context) + } + } + impl McpServerConnect for RmcpServer where Counterpart: Role, @@ -95,13 +126,34 @@ pub trait McpServerExt { } } - McpServer::new( - RmcpServer { - name: name.to_string(), - new_fn, - }, - NullRun, - ) + #[cfg(feature = "unstable_mcp_over_acp")] + { + // Feature unification must not construct an unused ACP service + // when this server is only used through its standalone adapter. + let new_fn = Arc::new(new_fn); + let shared = SharedRmcp { + new_fn: new_fn.clone(), + service: std::sync::OnceLock::new(), + }; + McpServer::new_service_with_standalone( + shared, + RmcpServer { + name: name.to_string(), + new_fn: move || new_fn(), + }, + NullRun, + ) + } + #[cfg(not(feature = "unstable_mcp_over_acp"))] + { + McpServer::new( + RmcpServer { + name: name.to_string(), + new_fn, + }, + NullRun, + ) + } } } @@ -130,10 +182,7 @@ where let byte_streams = ByteStreams::new(mcp_client_write.compat_write(), mcp_client_read.compat()); - // Spawn task to connect byte_streams to the provided client - drop(ConnectTo::::connect_to(byte_streams, client).await); - - Ok(()) + ConnectTo::::connect_to(byte_streams, client).await }; let bytes_to_rmcp = async { diff --git a/src/agent-client-protocol-rmcp/src/native.rs b/src/agent-client-protocol-rmcp/src/native.rs new file mode 100644 index 00000000..6c5cc9ba --- /dev/null +++ b/src/agent-client-protocol-rmcp/src/native.rs @@ -0,0 +1,239 @@ +//! Direct, request-scoped rmcp transport for ACP (no byte-stream emulation). + +use std::{ + future::Future, + sync::{Arc, Mutex}, +}; + +use acp::{ + Role, + mcp_server::{MCP_BACKEND_FAILURE, McpOutcome, McpRequest, McpRequestContext}, +}; +use agent_client_protocol as acp; +use futures::{ + channel::oneshot, + future::{BoxFuture, Either}, +}; +use rmcp::{ + RoleServer, Service, + model::ClientJsonRpcMessage, + service::{self, NotificationContext, RequestContext}, + transport::OneshotTransport, +}; +use tokio_util::sync::CancellationToken; + +/// Reuses the same application service but owns every handler future and its +/// cancellation on this one operation. +struct OperationService { + app: Arc, + cancel: CancellationToken, + completions: Arc>>>, +} + +impl> Service for OperationService { + fn handle_request( + &self, + request: ::PeerReq, + context: RequestContext, + ) -> impl Future::Resp, rmcp::ErrorData>> + + Send + + '_ { + let (done_tx, done_rx) = oneshot::channel(); + self.completions + .lock() + .expect("MCP operation poisoned") + .push(done_rx); + let cancel = self.cancel.clone(); + async move { + let result = tokio::select! { + biased; + () = cancel.cancelled() => Err(rmcp::ErrorData::internal_error("operation cancelled", None)), + result = self.app.handle_request(request, context) => result, + }; + let _sent = done_tx.send(()); + result + } + } + + fn handle_notification( + &self, + notification: ::PeerNot, + context: NotificationContext, + ) -> impl Future> + Send + '_ { + let (done_tx, done_rx) = oneshot::channel(); + self.completions + .lock() + .expect("MCP operation poisoned") + .push(done_rx); + let cancel = self.cancel.clone(); + async move { + let result = tokio::select! { + biased; + () = cancel.cancelled() => Ok(()), + result = self.app.handle_notification(notification, context) => result, + }; + let _sent = done_tx.send(()); + result + } + } + + fn get_info(&self) -> ::Info { + self.app.get_info() + } + + fn supported_protocol_versions( + &self, + ) -> std::borrow::Cow<'static, [rmcp::model::ProtocolVersion]> { + self.app.supported_protocol_versions() + } +} + +/// Execute one request against shared rmcp application state. Neither the +/// client transport nor the rmcp server's actor is allowed to escape this call. +pub(crate) fn execute( + app: Arc, + request: McpRequest, + context: McpRequestContext, +) -> BoxFuture<'static, Result> +where + R: Role, + S: Service, +{ + Box::pin(async move { + let id = context.request_id().0.to_string(); + let raw = serde_json::json!({ + "jsonrpc": "2.0", + "id": id, + "method": request.method, + "params": request.params, + }); + let inbound: ClientJsonRpcMessage = match serde_json::from_value(raw) { + Ok(request) => request, + Err(error) => { + return Ok(McpOutcome::Error( + acp::schema::v1::McpError::new( + if error.to_string().contains("unknown variant") { + -32601 + } else { + -32602 + }, + "Invalid MCP request", + ) + .data(serde_json::Value::String(error.to_string())), + )); + } + }; + let (transport, mut output) = OneshotTransport::::new(inbound); + let cancel = CancellationToken::new(); + let completions = Arc::new(Mutex::new(Vec::new())); + let handler = OperationService { + app, + cancel: cancel.clone(), + completions: completions.clone(), + }; + let mut running = service::serve_directly_with_ct(handler, transport, None, cancel.clone()); + let operation = async { + while let Some(outbound) = output.recv().await { + let value = serde_json::to_value(outbound).map_err(|error| { + acp::Error::new( + acp::mcp_server::MCP_BACKEND_FAILURE, + format!("cannot serialize MCP output: {error}"), + ) + })?; + match value { + serde_json::Value::Object(mut object) if object.contains_key("method") => { + if object.contains_key("id") { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "reverse MCP requests are not supported", + )); + } + let method = object + .remove("method") + .and_then(|v| v.as_str().map(str::to_owned)) + .ok_or_else(|| { + acp::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend notification has no method", + ) + })?; + let params = match object.remove("params") { + None | Some(serde_json::Value::Null) => None, + Some(serde_json::Value::Object(params)) => Some(params), + _ => { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend notification parameters must be an object", + )); + } + }; + context.send_notification(method, params).await?; + } + serde_json::Value::Object(mut object) if object.contains_key("result") => { + if object.get("id").and_then(serde_json::Value::as_str) != Some(id.as_str()) + { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "MCP response ID mismatch", + )); + } + return Ok(McpOutcome::Result( + object.remove("result").expect("checked result"), + )); + } + serde_json::Value::Object(mut object) if object.contains_key("error") => { + if object.get("id").and_then(serde_json::Value::as_str) != Some(id.as_str()) + { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "MCP error ID mismatch", + )); + } + let error = + serde_json::from_value(object.remove("error").expect("checked error")) + .map_err(|error| { + acp::Error::new( + acp::mcp_server::MCP_BACKEND_FAILURE, + format!("invalid MCP error from backend: {error}"), + ) + })?; + return Ok(McpOutcome::Error(error)); + } + _ => { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "unexpected MCP output", + )); + } + } + } + Err(acp::Error::new( + acp::mcp_server::MCP_BACKEND_FAILURE, + "MCP backend closed without a response", + )) + }; + let cancelled = async { + let acp = context.cancellation().cancelled(); + let operation = context.operation_cancellation().cancelled(); + futures::pin_mut!(acp, operation); + let _reason = futures::future::select(acp, operation).await; + }; + let result = match futures::future::select(Box::pin(operation), Box::pin(cancelled)).await { + Either::Left((result, _)) => result, + Either::Right(((), _)) => Err(acp::Error::request_cancelled()), + }; + cancel.cancel(); + let closed = running.close().await; + let handlers = std::mem::take(&mut *completions.lock().expect("MCP operation poisoned")); + for completion in handlers { + let _finished = completion.await; + } + closed.map_err(|error| { + acp::Error::new( + acp::mcp_server::MCP_BACKEND_FAILURE, + format!("MCP backend cleanup failed: {error}"), + ) + })?; + result + }) +} diff --git a/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs b/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs index 6ce95acd..fd3d5ac0 100644 --- a/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs +++ b/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs @@ -2,6 +2,7 @@ #![cfg(all(feature = "unstable_protocol_v2", feature = "unstable_mcp_over_acp"))] use std::{ + collections::HashMap, future::Future, sync::{Arc, Mutex}, time::Duration, @@ -47,7 +48,16 @@ async fn message( ) .block_task() .await?; - serde_json::from_str(response.0.get()).map_err(Error::into_internal_error) + match response { + v2::MessageMcpResponse::Result { result, .. } => Ok(result), + v2::MessageMcpResponse::Error { error, .. } => Err(Error::new(error.code, error.message) + .data(match error.data { + agent_client_protocol::schema::MaybeUndefined::Value(value) => Some(value), + agent_client_protocol::schema::MaybeUndefined::Null => Some(Value::Null), + agent_client_protocol::schema::MaybeUndefined::Undefined => None, + })), + _ => Err(Error::internal_error().data("unexpected MCP carrier outcome")), + } } struct DropSignal(Arc>>>); @@ -63,6 +73,7 @@ struct Service { _drop: DropSignal, started: Arc>>>, stopped: Arc>>>, + pending: Arc, oneshot::Sender<()>)>>>, } impl ServerHandler for Service { fn get_info(&self) -> ServerConfig { @@ -78,30 +89,56 @@ impl ServerHandler for Service { request: CallToolRequestParams, cx: RequestContext, ) -> impl Future> + Send { - std::future::ready(match request.name.as_ref() { - "retry" if request.request_state.is_none() => { - let inputs = serde_json::from_value(json!({"confirmation": { - "method": "elicitation/create", "params": {"mode": "form", - "message": "Confirm", "requestedSchema": {"type": "object", - "properties": {"approved": {"type": "boolean"}}}} - }})) - .expect("valid elicitation"); - Ok(InputRequiredResult::new(Some(inputs), Some("retry-state".into())).into()) + let pending = if request.name.as_ref() == "hang" { + let probe = request + .arguments + .as_ref() + .and_then(|args| args.get("probe")) + .and_then(Value::as_str) + .expect("pending tool requires a named probe"); + Some( + self.pending + .lock() + .unwrap() + .remove(probe) + .expect("distinct operation probe"), + ) + } else { + None + }; + async move { + if let Some((started, dropped)) = pending { + let _drop = DropSignal(Arc::new(Mutex::new(Some(dropped)))); + let _started = started.send(()); + // Deliberately ignore rmcp RequestContext::ct: the adapter must + // drop this future on outer cancellation and join its cleanup. + std::future::pending::<()>().await; } - "retry" if request.request_state.as_deref() == Some("retry-state") => Ok( - CallToolResult::structured(json!({"marker": cx.meta.get("example/marker"), + match request.name.as_ref() { + "retry" if request.request_state.is_none() => { + let inputs = serde_json::from_value(json!({"confirmation": { + "method": "elicitation/create", "params": {"mode": "form", + "message": "Confirm", "requestedSchema": {"type": "object", + "properties": {"approved": {"type": "boolean"}}}} + }})) + .expect("valid elicitation"); + Ok(InputRequiredResult::new(Some(inputs), Some("retry-state".into())).into()) + } + "retry" if request.request_state.as_deref() == Some("retry-state") => Ok( + CallToolResult::structured(json!({"marker": cx.meta.get("example/marker"), "responses": request.input_responses})) - .into(), - ), - "echo" => Ok(CallToolResult::structured( - json!({"marker": cx.meta.get("example/marker")}), - ) - .into()), - _ => Err(ErrorData::invalid_params( - "unknown tool or state", - Some(json!({"source": "rmcp"})), - )), - }) + .into(), + ), + "echo" => Ok(CallToolResult::structured( + json!({"marker": cx.meta.get("example/marker")}), + ) + .into()), + _ => Err(ErrorData::invalid_params( + "unknown tool or state", + Some(json!({"source": "rmcp"})), + )), + } + } } fn accepted_subscription_filter( &self, @@ -128,7 +165,7 @@ async fn exercise( server: v2::McpServerAcpId, started: oneshot::Receiver<()>, stopped: oneshot::Receiver<()>, - dropped: oneshot::Receiver<()>, + pending: Vec<(String, oneshot::Receiver<()>, oneshot::Receiver<()>)>, ) -> Result { let direct = message( &cx, @@ -215,7 +252,35 @@ async fn exercise( assert_eq!(parallel["structuredContent"]["marker"], "parallel"); subscription.cancel()?; stopped.await.map_err(Error::into_internal_error)?; - dropped.await.map_err(Error::into_internal_error)?; + for (index, (probe, started, dropped)) in pending.into_iter().enumerate() { + let id = format!("hang-{index}"); + let mut params = json!({"name": "hang", "arguments": {"probe": probe}}); + params["_meta"] = meta(&id); + let request = cx.send_request( + v2::MessageMcpRequest::new(server.clone(), id.clone(), "tools/call") + .params(params.as_object().expect("object params").clone()), + ); + started.await.map_err(Error::into_internal_error)?; + request.cancel()?; + // This is a distinct operation-local future, not the shared service's + // destructor. Cleanup must precede the cancellation response. + dropped.await.map_err(Error::into_internal_error)?; + let error = request + .block_task() + .await + .expect_err("cancelled MCP request"); + assert_eq!(i32::from(error.code), -32800); + let healthy = message( + &cx, + &server, + &format!("healthy-{index}"), + "tools/call", + json!({"name": "echo", "arguments": {}}), + &id, + ) + .await?; + assert_eq!(healthy["structuredContent"]["marker"], id); + } Ok(server) } @@ -226,7 +291,21 @@ async fn native_acp_stateless_rmcp_lifecycle() -> Result<(), Error> { let (stop_tx, stop_rx) = oneshot::channel(); let (drop_tx, drop_rx) = oneshot::channel(); let (result_tx, result_rx) = oneshot::channel(); - let invocation = Arc::new(Mutex::new(Some((start_rx, stop_rx, drop_rx, result_tx)))); + let mut pending_checks = Vec::new(); + let mut pending_handlers = HashMap::new(); + for name in ["first", "second"] { + let (started_tx, started_rx) = oneshot::channel(); + let (dropped_tx, dropped_rx) = oneshot::channel(); + pending_handlers.insert(name.to_owned(), (started_tx, dropped_tx)); + pending_checks.push((name.to_owned(), started_rx, dropped_rx)); + } + let pending_handlers = Arc::new(Mutex::new(pending_handlers)); + let invocation = Arc::new(Mutex::new(Some(( + start_rx, + stop_rx, + pending_checks, + result_tx, + )))); let (notifications_tx, mut notifications_rx) = mpsc::unbounded_channel(); let started = Arc::new(Mutex::new(Some(start_tx))); let stopped = Arc::new(Mutex::new(Some(stop_tx))); @@ -263,11 +342,12 @@ async fn native_acp_stateless_rmcp_lifecycle() -> Result<(), Error> { } other => panic!("unexpected declarations: {other:?}"), }; - let (start_rx, stop_rx, drop_rx, result_tx) = + let (start_rx, stop_rx, pending_checks, result_tx) = invocation.lock().unwrap().take().expect("one session"); let call_cx = cx.clone(); cx.spawn(async move { - let result = exercise(call_cx, server, start_rx, stop_rx, drop_rx).await; + let result = + exercise(call_cx, server, start_rx, stop_rx, pending_checks).await; drop(result_tx.send(result)); Ok(()) })?; @@ -287,16 +367,25 @@ async fn native_acp_stateless_rmcp_lifecycle() -> Result<(), Error> { agent_client_protocol::on_receive_notification!(), ); - Client.v2().connect_with(agent, async move |cx| { + let result = Client.v2().connect_with(agent, async move |cx| { cx.send_request(v2::InitializeRequest::new(ProtocolVersion::V2, v2::Implementation::new("native-rmcp-client", "1"))).block_task().await?; - let server = McpServer::::from_rmcp("real-rmcp", move || Service { - _drop: DropSignal(dropped.clone()), - started: started.clone(), stopped: stopped.clone(), + let created = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let factory_calls = created.clone(); + let server = McpServer::::from_rmcp("real-rmcp", move || { + factory_calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Service { + _drop: DropSignal(dropped.clone()), + started: started.clone(), stopped: stopped.clone(), + pending: pending_handlers.clone(), + } }); + assert_eq!(created.load(std::sync::atomic::Ordering::SeqCst), 0); cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) .with_mcp_server(server)?.start_session().block_task().await?; let server_id = result_rx.await.map_err(Error::into_internal_error)??; + assert_eq!(created.load(std::sync::atomic::Ordering::SeqCst), 1, + "independent native operations share one application service"); let acknowledgment = notifications_rx.recv().await.expect("acknowledgment"); let update = notifications_rx.recv().await.expect("filtered update"); assert_eq!(acknowledgment.method, "notifications/subscriptions/acknowledged"); @@ -313,7 +402,9 @@ async fn native_acp_stateless_rmcp_lifecycle() -> Result<(), Error> { ["io.modelcontextprotocol/subscriptionId"], json!("listen-1")); } Ok(()) - }).await + }).await; + drop_rx.await.map_err(Error::into_internal_error)?; + result }) .await .expect("native ACP/rmcp operation or cleanup timed out") diff --git a/src/agent-client-protocol/Cargo.toml b/src/agent-client-protocol/Cargo.toml index 4939a6e8..61f79dc3 100644 --- a/src/agent-client-protocol/Cargo.toml +++ b/src/agent-client-protocol/Cargo.toml @@ -62,6 +62,7 @@ wasm_js = ["uuid/js"] [dependencies] agent-client-protocol-schema.workspace = true agent-client-protocol-derive.workspace = true +async-channel.workspace = true futures.workspace = true futures-concurrency.workspace = true rustc-hash.workspace = true diff --git a/src/agent-client-protocol/examples/v2_session_coordination/tests.rs b/src/agent-client-protocol/examples/v2_session_coordination/tests.rs index a8a8aea7..68619f0b 100644 --- a/src/agent-client-protocol/examples/v2_session_coordination/tests.rs +++ b/src/agent-client-protocol/examples/v2_session_coordination/tests.rs @@ -1,7 +1,8 @@ use std::{future::Future, time::Duration}; use agent_client_protocol::{ - Channel, RawJsonRpcMessage, TransportBatch, TransportFrame, schema::v1::RequestId, + BudgetedFrame, Channel, RawJsonRpcMessage, TransportBatch, TransportFrame, + schema::v1::RequestId, }; use serde_json::{Value, json}; @@ -15,7 +16,7 @@ struct Peer(Channel); impl Peer { async fn request(&mut self, method: &str, session: Option<&str>) -> RequestId { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - self.0.rx.next().await + self.0.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected {method}"); }; @@ -34,7 +35,7 @@ impl Peer { fn respond(&self, id: RequestId, result: Result) { self.0 .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .try_send(TransportFrame::Single(RawJsonRpcMessage::response( id, result, ))) .unwrap(); @@ -43,7 +44,7 @@ impl Peer { fn replay_and_respond(&self, id: RequestId, session: &str, text: &str) { self.0 .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ update(session, text), RawJsonRpcMessage::response(id, Ok(json!({}))), @@ -83,7 +84,7 @@ impl Peer { while let Some(frame) = self.0.rx.next().await { assert!( matches!( - frame, + frame.frame(), TransportFrame::Single(RawJsonRpcMessage::Notification(_)) ), "unexpected request during shutdown: {frame:?}" @@ -191,7 +192,8 @@ async fn concurrent_loaders_share_one_resume_and_projection() { abandon.await.unwrap(); peer.0 .tx - .unbounded_send(TransportFrame::Single(update(SESSION, "hello"))) + .send_frame(TransportFrame::Single(update(SESSION, "hello"))) + .await .unwrap(); peer.respond(resume, Ok(response)); // A second resume or an early/duplicate close fails this script. @@ -228,13 +230,14 @@ async fn abandoned_resume_is_drained_and_closed_before_fresh_replay() { // Pre-close traffic must also drain before installing a new recipient. peer.0 .tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages([ update(SESSION, "closing"), RawJsonRpcMessage::response(close, Ok(json!({}))), ]) .unwrap(), )) + .await .unwrap(); let fresh = peer.request("session/resume", Some(SESSION)).await; peer.replay_and_respond(fresh, SESSION, "hello"); @@ -277,13 +280,14 @@ async fn delayed_close_blocks_only_its_session() { release_close.await.unwrap(); peer.0 .tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages([ update(OTHER, "+live"), RawJsonRpcMessage::response(close, Ok(json!({}))), ]) .unwrap(), )) + .await .unwrap(); let fresh = peer.request("session/resume", Some(SESSION)).await; peer.replay_and_respond(fresh, SESSION, "fresh"); @@ -370,7 +374,7 @@ async fn disconnect_during_cleanup_fails_waiting_reopen() { // A replacement resume must never have been published. while let Some(frame) = peer.0.rx.next().await { assert!(matches!( - frame, + frame.into_frame(), TransportFrame::Single(RawJsonRpcMessage::Notification(_)) )); } @@ -511,7 +515,7 @@ async fn eof_follows_received_replay_and_response_but_fails_unanswered_loads() { drop(peer.0.tx); while let Some(frame) = peer.0.rx.next().await { assert!(matches!( - frame, + frame.into_frame(), TransportFrame::Single(RawJsonRpcMessage::Notification(_)) )); } diff --git a/src/agent-client-protocol/src/acp_agent.rs b/src/agent-client-protocol/src/acp_agent.rs index 038fb388..9c103b04 100644 --- a/src/agent-client-protocol/src/acp_agent.rs +++ b/src/agent-client-protocol/src/acp_agent.rs @@ -1344,7 +1344,7 @@ mod tests { #[cfg(unix)] async fn reported_descendant_pid( - connection: &mut futures::future::BoxFuture<'static, Result<(), crate::Error>>, + connection: &mut (impl Future> + Unpin), pid_rx: &mut tokio::sync::mpsc::UnboundedReceiver, ) -> rustix::process::Pid { tokio::time::timeout(std::time::Duration::from_secs(5), async { @@ -1417,7 +1417,7 @@ mod tests { Ok(serde_json::json!({ "payload": "x".repeat(4 * 1024 * 1024) })), ); outgoing - .unbounded_send(crate::TransportFrame::Single(response)) + .try_send(crate::TransportFrame::Single(response)) .expect("response should be accepted before the connection starts"); outgoing.close_channel(); diff --git a/src/agent-client-protocol/src/component.rs b/src/agent-client-protocol/src/component.rs index ba8b9297..75b78c4a 100644 --- a/src/agent-client-protocol/src/component.rs +++ b/src/agent-client-protocol/src/component.rs @@ -27,10 +27,61 @@ //! ``` use futures::future::BoxFuture; -use std::{fmt::Debug, future::Future, marker::PhantomData}; +use std::{ + fmt::Debug, + future::Future, + marker::PhantomData, + pin::Pin, + task::{Context, Poll}, +}; use crate::{Channel, Result, role::Role}; +/// Connection work owned by a component, or a passive endpoint with no driver. +/// +/// Both can be awaited, but successful completion of a passive driver says +/// nothing about endpoint lifetime. Bridges must continue copying both halves +/// until they close. An active driver owns the component's completion signal. +pub struct ConnectionDriver(Option>>); + +impl ConnectionDriver { + /// Wrap work that owns a component's connection lifetime. + pub fn new(future: impl Future> + Send + 'static) -> Self { + Self(Some(Box::pin(future))) + } + + /// An endpoint whose I/O is driven elsewhere, such as an existing Channel. + #[must_use] + pub fn passive() -> Self { + Self(None) + } + + /// Whether completion is a no-op rather than an owned lifetime signal. + #[must_use] + pub fn is_passive(&self) -> bool { + self.0.is_none() + } +} + +impl Future for ConnectionDriver { + type Output = Result<()>; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + match self.0.as_mut() { + Some(future) => future.as_mut().poll(cx), + None => Poll::Ready(Ok(())), + } + } +} + +impl Debug for ConnectionDriver { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ConnectionDriver") + .field("passive", &self.is_passive()) + .finish() + } +} + /// A component that can exchange JSON-RPC messages to an endpoint playing the role `R` /// (e.g., an ACP [`Agent`](`crate::role::acp::Agent`) or an MCP [`Server`](`crate::role::mcp::Server`)). /// @@ -137,7 +188,8 @@ pub trait ConnectTo: Send + 'static { /// /// This method returns: /// - A `Channel` that can be used to communicate with this component - /// - A `BoxFuture` that drives the component's connection logic + /// - A [`ConnectionDriver`] that drives the component's connection logic, + /// or explicitly identifies an endpoint driven elsewhere /// /// The default implementation creates an intermediate channel pair and calls `connect_to` /// on one endpoint while returning the other endpoint for the caller to use. @@ -146,14 +198,14 @@ pub trait ConnectTo: Send + 'static { /// /// # Returns /// - /// A tuple of `(Channel, BoxFuture)` where the channel is for the caller to use + /// A tuple of `(Channel, ConnectionDriver)` where the channel is for the caller to use /// and the future must be polled to drive the connection. - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<()>>) + fn into_channel_and_future(self) -> (Channel, ConnectionDriver) where Self: Sized, { let (channel_a, channel_b) = Channel::duplex(); - let future = Box::pin(self.connect_to(channel_b)); + let future = ConnectionDriver::new(self.connect_to(channel_b)); (channel_a, future) } } @@ -171,8 +223,7 @@ trait ErasedConnectTo: Send { client: Box>, ) -> BoxFuture<'static, Result<()>>; - fn into_channel_and_future_erased(self: Box) - -> (Channel, BoxFuture<'static, Result<()>>); + fn into_channel_and_future_erased(self: Box) -> (Channel, ConnectionDriver); } /// Blanket implementation: any `ConnectTo` can be type-erased. @@ -195,9 +246,7 @@ impl, R: Role> ErasedConnectTo for C { }) } - fn into_channel_and_future_erased( - self: Box, - ) -> (Channel, BoxFuture<'static, Result<()>>) { + fn into_channel_and_future_erased(self: Box) -> (Channel, ConnectionDriver) { (*self).into_channel_and_future() } } @@ -251,7 +300,7 @@ impl ConnectTo for DynConnectTo { .await } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<()>>) { + fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { self.inner.into_channel_and_future_erased() } } diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 699ae9cf..8e22600a 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -18,13 +18,15 @@ use std::sync::{ Arc, Mutex, Weak, atomic::{AtomicBool, Ordering}, }; +use std::task::{Context, Poll, Waker}; use uuid::Uuid; use futures::FutureExt; -use futures::channel::{mpsc, oneshot}; +use futures::channel::oneshot; use futures::future::{self, BoxFuture, Either}; -use futures::{AsyncRead, AsyncWrite, StreamExt}; +use futures::{AsyncRead, AsyncWrite, Sink, SinkExt, StreamExt}; +mod admission; pub(crate) mod close; mod dynamic_handler; pub(crate) mod handlers; @@ -87,6 +89,31 @@ pub enum TransportFrame { Batch(TransportBatch), } +/// Finite transport and runtime admission limits. The byte budget is shared +/// across both directions of one in-memory duplex. +#[derive(Clone, Copy, Debug)] +pub struct ConnectionLimits { + /// Maximum UTF-8 bytes in one JSON-RPC frame. + pub max_frame_bytes: usize, + /// Shared serialized-payload budget, including queued frames and runtime + /// messages. One maximum frame's worth is reserved for responses/cancellation. + pub max_queued_bytes: usize, + /// Per-queue item limit and runtime admission limit for pending requests, + /// running tasks, dynamic handlers, and deferred dispatch. Values below one + /// are treated as one. Byte capacity is enforced separately. + pub max_queued_frames: usize, +} + +impl Default for ConnectionLimits { + fn default() -> Self { + Self { + max_frame_bytes: transport_actor::MAX_FRAME_BYTES, + max_queued_bytes: 64 * 1024 * 1024, + max_queued_frames: admission::QUEUE_CAPACITY, + } + } +} + /// A structurally non-empty JSON-RPC batch retained across framed relays. #[derive(Clone, Debug)] pub struct TransportBatch { @@ -236,6 +263,48 @@ impl Serialize for TransportBatch { } impl TransportFrame { + fn is_control(&self) -> bool { + fn message_is_control(message: &RawJsonRpcMessage) -> bool { + match message { + RawJsonRpcMessage::Response(_) => true, + RawJsonRpcMessage::Notification(notification) => { + if matches!( + notification.method.as_ref(), + "$/cancel_request" | "$/cancelRequest" + ) { + return true; + } + if !crate::schema::SuccessorMessage::::matches_method( + ¬ification.method, + ) { + return false; + } + let Some(RawJsonRpcParams::Object(envelope)) = ¬ification.params else { + return false; + }; + let Some(method) = envelope.get("method").and_then(serde_json::Value::as_str) + else { + return false; + }; + let (method, _) = peel_successor_envelopes( + method, + envelope.get("params").unwrap_or(&serde_json::Value::Null), + ); + matches!(method, "$/cancel_request" | "$/cancelRequest") + } + RawJsonRpcMessage::Request(_) => false, + } + } + match self { + Self::Single(message) => message_is_control(message), + Self::Batch(batch) => batch.entries().all(|entry| match entry { + TransportBatchEntry::Message(message) => message_is_control(message), + TransportBatchEntry::Malformed { .. } => false, + }), + Self::Malformed { .. } => false, + } + } + fn inspect_messages( &self, observer: &mut impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error>, @@ -1898,15 +1967,22 @@ impl< context: _, } = self; - let (outgoing_tx, outgoing_rx) = mpsc::unbounded(); - let (new_task_tx, new_task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, dynamic_handler_rx) = mpsc::unbounded(); - let pending_replies = PendingReplies::default(); - // Convert transport into server - this returns a channel for us to use // and a future that runs the transport. let transport_component = crate::DynConnectTo::new(transport); let (transport_channel, transport_future) = transport_component.into_channel_and_future(); + let limits = transport_channel.tx.admission().limits(); + let (outgoing_tx, outgoing_rx) = admission::budgeted_channel( + transport_channel.tx.admission(), + OutgoingMessage::charged_bytes, + OutgoingMessage::with_permit, + OutgoingMessage::is_control, + OutgoingMessage::is_urgent, + ); + let (new_task_tx, new_task_rx) = admission::channel_with_capacity(limits.max_queued_frames); + let (dynamic_handler_tx, dynamic_handler_rx) = + admission::channel_with_capacity(limits.max_queued_frames); + let pending_replies = PendingReplies::with_capacity(limits.max_queued_frames); let (transport_completion_tx, transport_completion_rx) = oneshot::channel(); let transport_completion = transport_completion_rx .map(|result| { @@ -1965,8 +2041,13 @@ impl< pending_replies, transport_outgoing_tx, protocol_compat, + connection.incoming_closed.clone(), + ), + task_actor::task_actor( + new_task_rx, + &connection, + limits.max_queued_frames ), - task_actor::task_actor(new_task_rx, &connection), runner.run_with_connection_to(connection.clone()), )?; Ok(()) @@ -1985,8 +2066,16 @@ impl< }; run_until_connection_close( - background, - main_fn(connection.clone()), + async { + let result = background.await; + connection.incoming_closed.request_shutdown(); + result + }, + async { + let result = main_fn(connection.clone()).await; + connection.incoming_closed.request_shutdown(); + result + }, connection.incoming_closed.clone(), ) .await @@ -2098,6 +2187,8 @@ pub(crate) struct ResponsePayload { /// the dispatch loop; ordinary blocking consumers, local error paths, and /// responses routed later do not. pub(crate) ack_tx: Option>, + /// Admission remains with an SDK-owned result until it is consumed or dropped. + retained_bytes: Option, } type ResponseRouteHook = @@ -2141,7 +2232,7 @@ impl std::fmt::Debug for ResponsePayload { f.debug_struct("ResponsePayload") .field("result", &self.result) .field("ack_tx", &self.ack_tx.as_ref().map(|_| "...")) - .finish() + .finish_non_exhaustive() } } @@ -2162,6 +2253,8 @@ impl ResponseOrdering { struct PendingReply { method: String, + /// The method and map key outlive the outgoing frame. + metadata_bytes: Option, role_id: RoleId, sender: oneshot::Sender, cancellation_disarm: SentRequestCancellationDisarm, @@ -2177,6 +2270,7 @@ impl PendingReply { .send(ResponsePayload { result: Err(error), ack_tx: None, + retained_bytes: self.metadata_bytes, }) .is_err() { @@ -2190,10 +2284,20 @@ impl PendingReply { } } -#[derive(Default)] struct PendingRepliesInner { incoming_closed: bool, replies: HashMap, + max_pending: usize, +} + +impl Default for PendingRepliesInner { + fn default() -> Self { + Self { + incoming_closed: false, + replies: HashMap::new(), + max_pending: admission::QUEUE_CAPACITY, + } + } } #[derive(Clone, Default)] @@ -2202,6 +2306,15 @@ struct PendingReplies { } impl PendingReplies { + fn with_capacity(max_pending: usize) -> Self { + Self { + inner: Arc::new(Mutex::new(PendingRepliesInner { + max_pending: max_pending.max(1), + ..PendingRepliesInner::default() + })), + } + } + fn registrar(&self) -> PendingRepliesRegistrar { PendingRepliesRegistrar { inner: Arc::downgrade(&self.inner), @@ -2224,6 +2337,36 @@ impl PendingReplies { .remove(id) } + fn mark_published(&self, id: &RequestId) -> bool { + let inner = self.inner.lock().expect("pending replies mutex poisoned"); + let Some(reply) = inner.replies.get(id) else { + return false; + }; + reply + .cancellation_disarm + .published + .store(true, Ordering::Release); + true + } + + /// Cancellation may bypass queued work, but must never reach the peer + /// before a request that we subsequently publish. Settle that case locally. + fn cancel_unpublished(&self, id: &RequestId) -> bool { + let reply = { + let mut inner = self.inner.lock().expect("pending replies mutex poisoned"); + if inner + .replies + .get(id) + .is_none_or(|reply| reply.cancellation_disarm.published.load(Ordering::Acquire)) + { + return false; + } + inner.replies.remove(id).expect("pending reply checked") + }; + reply.fail(crate::Error::request_cancelled()); + true + } + /// Atomically reject new subscriptions and fail every existing one. fn close_incoming(&self) -> usize { let replies = { @@ -2274,6 +2417,12 @@ impl PendingRepliesRegistrar { let mut inner = inner.lock().expect("pending replies mutex poisoned"); if inner.incoming_closed { Err(reply) + } else if !inner.replies.contains_key(&id) && inner.replies.len() >= inner.max_pending { + drop(inner); + reply.fail(crate::util::internal_error( + "pending request capacity exceeded", + )); + return false; } else { Ok(inner.replies.insert(id, reply)) } @@ -2304,6 +2453,17 @@ impl PendingRepliesRegistrar { .replies .remove(id) } + + fn discard_abandoned(&self, id: &RequestId) -> Option { + let inner = self.inner.upgrade()?; + let mut inner = inner.lock().expect("pending replies mutex poisoned"); + // Framework response hooks own cleanup even when their consumer drops. + // Keep their bounded registration until the reply arrives or EOF fails it. + if inner.replies.get(id)?.response_route_hook.is_some() { + return None; + } + inner.replies.remove(id) + } } impl Debug for PendingRepliesRegistrar { @@ -2653,6 +2813,14 @@ fn peel_successor_envelopes<'message>( (method, params) } +fn outgoing_cancellation_id(message: &UntypedMessage) -> Option { + let (method, params) = peel_successor_envelopes(&message.method, &message.params); + if !matches!(method, "$/cancel_request" | "$/cancelRequest") { + return None; + } + serde_json::from_value(params.get("requestId")?.clone()).ok() +} + /// Whether a notification is a `$/cancel_request`, even when it is still /// wrapped in `_proxy/successor` envelopes. /// @@ -2718,6 +2886,7 @@ impl ResponseDestination { remaining: slot_count, responses: (0..slot_count).map(|_| None).collect(), abandoned: (0..slot_count).map(|_| None).collect(), + permits: (0..slot_count).map(|_| None).collect(), active_handler_attempts: (0..slot_count).map(|_| 0).collect(), dispatch_complete: false, emitted: false, @@ -2737,17 +2906,29 @@ impl ResponseDestination { ) } - fn complete(self, response: RawJsonRpcMessage) -> Option { + fn complete_admitted( + self, + response: RawJsonRpcMessage, + permit: Option, + ) -> Option<(TransportFrame, Option)> { match self { - Self::Individual(slot) => slot.complete(response), - Self::Batch(slot) => slot.complete(response).map(batch_response_frame), + Self::Individual(slot) => slot.complete(response).map(|frame| (frame, permit)), + Self::Batch(slot) => slot + .complete_admitted(response, permit) + .map(batch_response_frame_admitted), } } - fn abandon(self, fallback: RawJsonRpcMessage) -> Option { + fn abandon_admitted( + self, + fallback: RawJsonRpcMessage, + permit: Option, + ) -> Option<(TransportFrame, Option)> { match self { Self::Individual(_) => None, - Self::Batch(slot) => slot.abandon(fallback).map(batch_response_frame), + Self::Batch(slot) => slot + .abandon_admitted(fallback, permit) + .map(batch_response_frame_admitted), } } @@ -2769,10 +2950,19 @@ impl ResponseDestination { }) } - fn finish_handler_attempt(self) -> Option { + fn finish_handler_attempt_admitted( + self, + permit: Option, + ) -> Option<(TransportFrame, Option)> { match self { Self::Individual(_) => None, - Self::Batch(slot) => slot.finish_handler_attempt().map(batch_response_frame), + Self::Batch(slot) => slot.finish_handler_attempt().map(|ready| { + let (frame, mut charge) = batch_response_frame_admitted(ready); + if let (Some(charge), Some(permit)) = (&mut charge, permit) { + charge.join(permit); + } + (frame, charge) + }), } } } @@ -2800,6 +2990,20 @@ fn batch_response_frame(responses: Vec) -> TransportFrame { ) } +fn batch_response_frame_admitted(ready: BatchReady) -> (TransportFrame, Option) { + let mut permits = ready.permits.into_iter(); + let mut charge = permits.next(); + for permit in permits { + charge.as_mut().expect("first permit exists").join(permit); + } + (batch_response_frame(ready.responses), charge) +} + +struct BatchReady { + responses: Vec, + permits: Vec, +} + #[derive(Clone)] struct BatchDispatchCompletion { state: Arc>, @@ -2814,7 +3018,10 @@ impl std::fmt::Debug for BatchDispatchCompletion { } impl BatchDispatchCompletion { - fn complete(self) -> Option { + fn complete_admitted( + self, + permit: Option, + ) -> Option<(TransportFrame, Option)> { let mut state = self .state .lock() @@ -2827,7 +3034,13 @@ impl BatchDispatchCompletion { for index in 0..state.responses.len() { promote_abandoned_response(&mut state, index); } - take_completed_batch(&mut state).map(batch_response_frame) + take_completed_batch(&mut state).map(|ready| { + let (frame, mut charge) = batch_response_frame_admitted(ready); + if let (Some(charge), Some(permit)) = (&mut charge, permit) { + charge.join(permit); + } + (frame, charge) + }) } } @@ -2841,14 +3054,14 @@ fn promote_abandoned_response(state: &mut BatchResponseState, index: usize) { } } -fn take_completed_batch(state: &mut BatchResponseState) -> Option> { +fn take_completed_batch(state: &mut BatchResponseState) -> Option { if !state.dispatch_complete || state.remaining != 0 || state.emitted { return None; } state.emitted = true; - Some( - state + Some(BatchReady { + responses: state .responses .iter_mut() .map(|response| { @@ -2857,7 +3070,8 @@ fn take_completed_batch(state: &mut BatchResponseState) -> Option Option> { + fn finish_handler_attempt(self) -> Option { let mut state = self .state .lock() @@ -2898,7 +3112,11 @@ impl BatchResponseSlot { take_completed_batch(&mut state) } - fn complete(self, response: RawJsonRpcMessage) -> Option> { + fn complete_admitted( + self, + response: RawJsonRpcMessage, + permit: Option, + ) -> Option { let mut state = self .state .lock() @@ -2924,11 +3142,16 @@ impl BatchResponseSlot { state.abandoned[self.index] = None; state.responses[self.index] = Some(response); + state.permits[self.index] = permit; state.remaining -= 1; take_completed_batch(&mut state) } - fn abandon(self, fallback: RawJsonRpcMessage) -> Option> { + fn abandon_admitted( + self, + fallback: RawJsonRpcMessage, + permit: Option, + ) -> Option { let mut state = self .state .lock() @@ -2950,6 +3173,7 @@ impl BatchResponseSlot { } else { state.abandoned[self.index] = Some(fallback); } + state.permits[self.index] = permit; take_completed_batch(&mut state) } } @@ -2958,6 +3182,7 @@ struct BatchResponseState { remaining: usize, responses: Vec>, abandoned: Vec>, + permits: Vec>, active_handler_attempts: Vec, dispatch_complete: bool, emitted: bool, @@ -2995,6 +3220,8 @@ struct ResponseReplyTarget { sender: Arc>>>, ordering: ResponseOrdering, dispatch: ResponseDispatch, + /// Keep the original frame admitted while a handler defers routing. + frame_bytes: Option, } impl ResponseReplyTarget { @@ -3013,8 +3240,39 @@ impl ResponseReplyTarget { return; }; + // A transformed result may be larger than the wire response. Each + // result (including each member of a batch) therefore needs its own + // charge; cloning the batch's frame permit does not charge each result. + // Never wait here: this router may hold the only permit whose release + // would make room. On rejection deliver a bounded error instead. + let (result, retained_bytes) = if let Some(frame) = self.frame_bytes { + let bytes = match &result { + Ok(value) => serde_json::to_vec(value).map(|json| json.len()), + Err(error) => serde_json::to_vec(error).map(|json| json.len()), + }; + match bytes.ok().and_then(|bytes| { + FrameAdmission(frame.inner.budget.clone()).try_reserve_bytes(bytes, true) + }) { + Some(permit) => (result, Some(permit)), + None => ( + Err(crate::util::internal_error( + "retained response byte capacity exceeded", + )), + Some(frame), + ), + } + } else { + (result, None) + }; let ack_tx = self.dispatch.acknowledgment(&self.ordering); - if sender.send(ResponsePayload { result, ack_tx }).is_err() { + if sender + .send(ResponsePayload { + result, + ack_tx, + retained_bytes, + }) + .is_err() + { tracing::debug!( method = %self.method, id = ?self.id, @@ -3064,7 +3322,7 @@ impl ResponseDispatch { enum HandlerErrorTarget { Request(RequestReplyTarget), - Response(ResponseReplyTarget), + Response(Box), } impl HandlerErrorTarget { @@ -3081,6 +3339,13 @@ impl HandlerErrorTarget { #[derive(Debug)] enum OutgoingMessage { + /// Retain application admission across queueing, readiness, conversion, and + /// transport publication. Legacy test-only queues can still carry bare messages. + Admitted { + message: Box, + permit: FramePermit, + }, + /// Close the outgoing application queue and acknowledge after every /// already-accepted message has entered the raw transport queue. CloseAfterDraining { done: oneshot::Sender<()> }, @@ -3147,6 +3412,99 @@ enum OutgoingMessage { }, } +impl OutgoingMessage { + fn charged_bytes(&self) -> Result { + // Include space for the JSON-RPC envelope and request ID. A transformed + // frame that exceeds this estimate must grow the *same* permit, never + // await an independent reservation while retaining the first. + const ENVELOPE: usize = 64; + let bytes = match self { + Self::Admitted { message, .. } => return message.charged_bytes(), + Self::Request { + id, + method, + untyped, + .. + } => { + serde_json::to_vec(&untyped.params) + .map_err(crate::Error::into_internal_error)? + .len() + + untyped.method.len() + + method.len() + + serde_json::to_vec(id) + .map_err(crate::Error::into_internal_error)? + .len() + } + Self::Notification { untyped } => { + serde_json::to_vec(&untyped.params) + .map_err(crate::Error::into_internal_error)? + .len() + + untyped.method.len() + } + Self::Response { + id, + method, + response, + .. + } => { + serde_json::to_vec(response) + .map_err(crate::Error::into_internal_error)? + .len() + + method.len() + + serde_json::to_vec(id) + .map_err(crate::Error::into_internal_error)? + .len() + } + Self::UncorrelatedErrorResponse { error, .. } => serde_json::to_vec(error) + .map_err(crate::Error::into_internal_error)? + .len(), + Self::AbandonedBatchResponse { id, method, .. } => { + method.len() + + serde_json::to_vec(id) + .map_err(crate::Error::into_internal_error)? + .len() + } + Self::CloseAfterDraining { .. } + | Self::BatchDispatchComplete { .. } + | Self::BatchHandlerAttemptComplete { .. } => 0, + }; + Ok(if bytes == 0 { + 1 + } else { + bytes.saturating_add(ENVELOPE) + }) + } + + fn is_control(&self) -> bool { + match self { + Self::Admitted { message, .. } => message.is_control(), + Self::Notification { untyped } => outgoing_cancellation_id(untyped).is_some(), + Self::CloseAfterDraining { .. } + | Self::BatchDispatchComplete { .. } + | Self::BatchHandlerAttemptComplete { .. } + | Self::Response { .. } + | Self::UncorrelatedErrorResponse { .. } + | Self::AbandonedBatchResponse { .. } => true, + Self::Request { .. } => false, + } + } + + fn is_urgent(&self) -> bool { + match self { + Self::Admitted { message, .. } => message.is_urgent(), + Self::Notification { untyped } => outgoing_cancellation_id(untyped).is_some(), + _ => false, + } + } + + fn with_permit(self, permit: FramePermit) -> Self { + Self::Admitted { + message: Box::new(self), + permit, + } + } +} + /// Return type from JrHandler; indicates whether the request was handled or not. #[must_use] #[derive(Debug)] @@ -3229,6 +3587,11 @@ impl V2ConnectionTo { self.inner.incoming_closed().await; } + /// Wait for EOF or connection termination, before close callbacks run. + pub async fn shutdown_requested(&self) { + self.inner.shutdown_requested().await; + } + /// Return whether clean incoming-EOF processing has completed. #[must_use] pub fn is_incoming_closed(&self) -> bool { @@ -3337,6 +3700,17 @@ impl V2ConnectionTo { self.inner.send_notification(notification) } + /// Await outbound capacity outside the dispatch loop. + pub async fn send_notification_async( + &self, + notification: N, + ) -> Result<(), crate::Error> + where + Counterpart: HasPeer, + { + self.inner.send_notification_async(notification).await + } + /// Send an outgoing notification to a specific peer. pub fn send_notification_to( &self, @@ -3349,6 +3723,20 @@ impl V2ConnectionTo { self.inner.send_notification_to(peer, notification) } + /// Await outbound capacity outside the dispatch loop. + pub async fn send_notification_to_async( + &self, + peer: Peer, + notification: N, + ) -> Result<(), crate::Error> + where + Counterpart: HasPeer, + { + self.inner + .send_notification_to_async(peer, notification) + .await + } + /// Send a `$/cancel_request` notification to the default counterpart peer. pub fn send_cancel_request( &self, @@ -3431,7 +3819,7 @@ pub struct ConnectionTo { counterpart: Counterpart, message_tx: OutgoingMessageTx, task_tx: TaskTx, - dynamic_handler_tx: mpsc::UnboundedSender>, + dynamic_handler_tx: admission::Sender>, transport_completion: SharedTransportCompletion, pending_replies: PendingRepliesRegistrar, #[cfg_attr( @@ -3457,23 +3845,45 @@ struct IncomingClosedState { closed: AtomicBool, signal_tx: Mutex>>, signal_rx: future::Shared>, + shutdown_tx: Mutex>>, + shutdown_rx: future::Shared>, } impl IncomingClosed { fn new() -> Self { let (signal_tx, signal_rx) = oneshot::channel(); + let (shutdown_tx, shutdown_rx) = oneshot::channel(); Self { state: Arc::new(IncomingClosedState { closing: AtomicBool::new(false), closed: AtomicBool::new(false), signal_tx: Mutex::new(Some(signal_tx)), signal_rx: signal_rx.map(|_| ()).boxed().shared(), + shutdown_tx: Mutex::new(Some(shutdown_tx)), + shutdown_rx: shutdown_rx.map(|_| ()).boxed().shared(), }), } } fn begin_close(&self) { self.state.closing.store(true, Ordering::Release); + self.request_shutdown(); + } + + fn request_shutdown(&self) { + if let Some(tx) = self + .state + .shutdown_tx + .lock() + .expect("shutdown mutex poisoned") + .take() + { + let _ = tx.send(()); + } + } + + async fn shutdown_requested(&self) { + self.state.shutdown_rx.clone().await; } fn finish_close(&self) { @@ -3583,9 +3993,9 @@ fn run_until_connection_close( impl ConnectionTo { fn new( counterpart: Counterpart, - message_tx: mpsc::UnboundedSender, - task_tx: mpsc::UnboundedSender, - dynamic_handler_tx: mpsc::UnboundedSender>, + message_tx: OutgoingMessageTx, + task_tx: TaskTx, + dynamic_handler_tx: admission::Sender>, transport_completion: SharedTransportCompletion, pending_replies: PendingRepliesRegistrar, protocol_mode: ProtocolMode, @@ -3624,6 +4034,12 @@ impl ConnectionTo { self.incoming_closed.closed().await; } + /// Resolves on transport EOF or local completion, before close callbacks + /// or outgoing drain. Cancel connection-owned work when this fires. + pub async fn shutdown_requested(&self) { + self.incoming_closed.shutdown_requested().await; + } + /// Return whether clean incoming-EOF processing has completed. /// /// This remains `false` while [`Builder::on_close`] callbacks are running. @@ -3636,10 +4052,11 @@ impl ConnectionTo { /// the protocol actor, and wait for the transport sink to finish them. async fn drain_outgoing(&self) -> Result<(), crate::Error> { let (done_tx, done_rx) = oneshot::channel(); - let marker_result = send_raw_message( - &self.message_tx, - OutgoingMessage::CloseAfterDraining { done: done_tx }, - ); + let marker_result = self + .message_tx + .send(OutgoingMessage::CloseAfterDraining { done: done_tx }) + .await + .map_err(crate::util::internal_error); let marker_result = match marker_result { Ok(()) => done_rx.await.map_err(|error| { crate::util::internal_error(format!( @@ -4094,13 +4511,18 @@ impl ConnectionTo { } let role_id = peer.role_id(); let remote_style = self.counterpart.remote_style(peer); - let cancellation = - SentRequestCancellation::new(self.message_tx.clone(), remote_style, id.clone()); + let cancellation = SentRequestCancellation::new( + self.message_tx.clone(), + self.pending_replies.clone(), + remote_style, + id.clone(), + ); if self.is_incoming_closing() { cancellation.disarm(); drop(response_tx.send(ResponsePayload { result: Err(incoming_transport_closed_error(&method)), ack_tx: None, + retained_bytes: None, })); return SentRequest::new( id, @@ -4115,12 +4537,47 @@ impl ConnectionTo { match request.to_untyped_message() { Ok(untyped) => { + // The queue's frame charge is released after transport publication, + // but the pending map retains its own copies of the method and ID. + // Charge those strings (plus a fixed entry allowance) separately. + let metadata_bytes = self.message_tx.byte_admission().and_then(|budget| { + budget.try_reserve_bytes( + method + .len() + .saturating_add(match &id { + RequestId::Str(value) => value.len(), + _ => 32, + }) + .saturating_add(64), + true, + ) + }); + if self.message_tx.byte_admission().is_some() && metadata_bytes.is_none() { + cancellation.disarm(); + drop(response_tx.send(ResponsePayload { + result: Err(crate::util::internal_error( + "pending request metadata byte capacity exceeded", + )), + ack_tx: None, + retained_bytes: None, + })); + return SentRequest::new( + id, + method.clone(), + self.task_tx.clone(), + response_rx, + cancellation, + response_ordering, + ) + .map(move |json| ::from_value(&method, json)); + } // Register before enqueueing so incoming EOF can fail every // observable request before close callbacks begin. The // outgoing actor checks that the registration still exists // before sending the request. let pending_reply = PendingReply { method: method.clone(), + metadata_bytes, role_id, sender: response_tx, cancellation_disarm: cancellation.disarm_handle(), @@ -4142,10 +4599,9 @@ impl ConnectionTo { if let Err(error) = self.message_tx.unbounded_send(message) { cancellation.disarm(); - - let OutgoingMessage::Request { id, method, .. } = error.into_inner() else { - unreachable!(); - }; + // A rejected queue item may be wrapped in Admitted. + // Drop it to release its admission before failing the waiter. + drop(error.into_inner()); if let Some(pending_reply) = self.pending_replies.remove(&id) { if self.is_incoming_closing() { @@ -4169,6 +4625,7 @@ impl ConnectionTo { "failed to create untyped request for `{method}`: {err}" ))), ack_tx: None, + retained_bytes: None, }) .unwrap(); } @@ -4212,6 +4669,18 @@ impl ConnectionTo { self.send_notification_to(self.counterpart.clone(), notification) } + /// Await outbound capacity for a producer outside ordered dispatch. + pub async fn send_notification_async( + &self, + notification: N, + ) -> Result<(), crate::Error> + where + Counterpart: HasPeer, + { + self.send_notification_to_async(self.counterpart.clone(), notification) + .await + } + /// Send an outgoing notification to a specific peer (no reply expected). /// /// The message will be transformed according to the [`HasPeer`](crate::role::HasPeer) @@ -4246,6 +4715,25 @@ impl ConnectionTo { ) } + /// Await outbound capacity for a producer outside ordered dispatch. + pub async fn send_notification_to_async( + &self, + peer: Peer, + notification: N, + ) -> Result<(), crate::Error> + where + Counterpart: HasPeer, + { + let remote_style = self.counterpart.remote_style(peer); + let transformed = remote_style.transform_outgoing_message(notification)?; + self.message_tx + .send(OutgoingMessage::Notification { + untyped: transformed, + }) + .await + .map_err(crate::util::internal_error) + } + /// Send a `$/cancel_request` notification for an arbitrary request ID to /// the default counterpart peer. /// @@ -4722,7 +5210,7 @@ pub struct ResponseRouter { send_fn: Box) -> Result<(), crate::Error> + Send>, /// Shared route used to deliver a dispatch-handler error to the same waiter. - reply_target: ResponseReplyTarget, + reply_target: Box, } impl std::fmt::Debug for ResponseRouter { @@ -4741,9 +5229,15 @@ impl ResponseRouter { /// When [`route_with_result`](Self::route_with_result) is called, the response is sent through the oneshot /// channel to the code that originally sent the request. If that receiver was /// dropped, the response is discarded because there is no local awaiter left. - fn new(id: RequestId, pending_reply: PendingReply, dispatch: ResponseDispatch) -> Self { + fn new( + id: RequestId, + pending_reply: PendingReply, + dispatch: ResponseDispatch, + frame_bytes: Option, + ) -> Self { let PendingReply { method, + metadata_bytes: _, role_id, sender, cancellation_disarm, @@ -4756,6 +5250,7 @@ impl ResponseRouter { sender: Arc::new(Mutex::new(Some(sender))), ordering, dispatch, + frame_bytes, }; let send_target = reply_target.clone(); // A response for the request reached this router, so the request is @@ -4779,7 +5274,7 @@ impl ResponseRouter { send_target.route(response); Ok(()) }), - reply_target, + reply_target: Box::new(reply_target), } } @@ -5474,12 +5969,14 @@ pub struct SentRequest { #[derive(Clone, Debug)] pub(crate) struct SentRequestCancellationDisarm { armed: Arc, + published: Arc, } impl SentRequestCancellationDisarm { fn new() -> Self { Self { armed: Arc::new(AtomicBool::new(true)), + published: Arc::new(AtomicBool::new(false)), } } @@ -5490,6 +5987,8 @@ impl SentRequestCancellationDisarm { struct SentRequestCancellation { message_tx: OutgoingMessageTx, + pending_replies: PendingRepliesRegistrar, + retain_pending_on_drop: AtomicBool, remote_style: crate::role::RemoteStyle, request_id: RequestId, disarm: SentRequestCancellationDisarm, @@ -5498,11 +5997,14 @@ struct SentRequestCancellation { impl SentRequestCancellation { fn new( message_tx: OutgoingMessageTx, + pending_replies: PendingRepliesRegistrar, remote_style: crate::role::RemoteStyle, request_id: RequestId, ) -> Self { Self { message_tx, + pending_replies, + retain_pending_on_drop: AtomicBool::new(false), remote_style, request_id, disarm: SentRequestCancellationDisarm::new(), @@ -5537,6 +6039,11 @@ impl Drop for SentRequestCancellation { if let Err(error) = self.send() { tracing::debug!(?error, "failed to auto-cancel dropped request"); } + // The receiver is gone now; waiting for a peer response would retain + // the pending method and map key without any possible consumer. + if !self.retain_pending_on_drop.load(Ordering::Acquire) { + self.pending_replies.discard_abandoned(&self.request_id); + } } } @@ -5616,7 +6123,7 @@ impl SentRequest { fn new( id: RequestId, method: String, - task_tx: mpsc::UnboundedSender, + task_tx: TaskTx, response_rx: oneshot::Receiver, cancellation: SentRequestCancellation, response_ordering: ResponseOrdering, @@ -5649,6 +6156,11 @@ impl SentRequest { /// handle while automatic cancellation is armed. pub fn detach(self) { self.cancellation.disarm(); + // A detached request must stay registered until it has been + // published: the outgoing actor skips unregistered requests. + self.cancellation + .retain_pending_on_drop + .store(true, Ordering::Release); } /// Send a `$/cancel_request` notification for this outgoing request. @@ -5870,7 +6382,11 @@ impl SentRequest { .await; match response { - Ok(ResponsePayload { result, ack_tx }) => { + Ok(ResponsePayload { + result, + ack_tx, + retained_bytes, + }) => { // Convert the result using to_result for Ok values let typed_result = match result { Ok(json_value) => to_result(json_value), @@ -5878,6 +6394,7 @@ impl SentRequest { }; let outcome = handle(Ok(typed_result)).await; + drop(retained_bytes); // Ack AFTER the handler completes - this is the key // difference from block_task. The dispatch loop waits for @@ -5974,6 +6491,7 @@ impl SentRequest { Ok(ResponsePayload { result: Ok(json_value), ack_tx, + retained_bytes: _, }) => { // Blocking consumers ack before converting or returning the // value, so dispatch can continue while the caller processes it. @@ -5988,6 +6506,7 @@ impl SentRequest { Ok(ResponsePayload { result: Err(err), ack_tx, + retained_bytes: _, }) => { if let Some(tx) = ack_tx { let _ = tx.send(()); @@ -6020,11 +6539,16 @@ impl SentRequest { .await; let (result, ack_tx) = match response { - Ok(ResponsePayload { result, ack_tx }) => { + Ok(ResponsePayload { + result, + ack_tx, + retained_bytes, + }) => { let typed_result = match result { Ok(json_value) => (self.to_result)(json_value), Err(error) => Err(error), }; + drop(retained_bytes); (typed_result, ack_tx) } Err(error) => ( @@ -6339,8 +6863,9 @@ where } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { - self.into_channel_transport() + fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { + let (channel, driver) = self.into_channel_transport(); + (channel, crate::ConnectionDriver::new(driver)) } } @@ -6404,11 +6929,9 @@ where impl futures::Sink + Send + 'static, impl futures::Stream> + Send + 'static, > { - use futures::AsyncBufReadExt; - use futures::io::BufReader; let Self { outgoing, incoming } = self; - let incoming_lines = Box::pin(BufReader::new(incoming).lines()); + let incoming_lines = Box::pin(transport_actor::bounded_lines(Box::pin(incoming))); let outgoing_lines = futures::sink::unfold(Box::pin(outgoing), async move |mut writer, line: String| { write_line(&mut writer, line).await?; @@ -6440,7 +6963,7 @@ where ConnectTo::::connect_to(self.into_lines(), client).await } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { + fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { ConnectTo::::into_channel_and_future(self.into_lines()) } } @@ -6470,122 +6993,1451 @@ where #[derive(Debug)] pub struct Channel { /// Receives frames from the counterpart. - pub rx: mpsc::UnboundedReceiver, + pub rx: FrameReceiver, /// Sends frames to the counterpart. - pub tx: mpsc::UnboundedSender, + pub tx: FrameSender, } -impl Channel { - /// Create a pair of connected channel endpoints. - /// - /// Frames sent through either endpoint are received by the other endpoint. - #[must_use] - pub fn duplex() -> (Self, Self) { - let (a_tx, b_rx) = mpsc::unbounded(); - let (b_tx, a_rx) = mpsc::unbounded(); +/// The byte charge for a frame. Clones refer to the same charge; releasing it +/// requires dropping *every* copy, including deferred dispatch/writer copies. +#[derive(Clone, Debug)] +pub struct FramePermit { + inner: Arc, + additional: Vec, +} - (Self { rx: a_rx, tx: a_tx }, Self { rx: b_rx, tx: b_tx }) +impl FramePermit { + /// Number of bytes held until every copy of this permit is dropped. + pub fn charged_bytes(&self) -> usize { + self.inner.bytes.load(Ordering::Acquire) + + self + .additional + .iter() + .map(Self::charged_bytes) + .sum::() } - /// Copy frames from `rx` to `tx` until the input closes. - /// - /// # Errors - /// - /// Returns an error if the receiving endpoint closes before the input. - pub(crate) async fn copy(mut self) -> Result<(), crate::Error> { - while let Some(frame) = self.rx.next().await { - self.tx - .unbounded_send(frame) - .map_err(crate::util::internal_error)?; - } - Ok(()) + fn join(&mut self, other: FramePermit) { + self.additional.push(other); } - /// Bridge two endpoints while inspecting every valid message. - /// - /// Observers are invoked in source order, including for each valid member of - /// a batch. The original frame is forwarded unchanged after inspection. - /// - /// # Errors - /// - /// Returns an observer error or an error if a destination closes before its - /// source. - pub async fn bridge_with_inspection( - left: Self, - right: Self, - mut left_to_right: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send, - mut right_to_left: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send, - ) -> Result<(), crate::Error> { - let Self { - rx: mut left_rx, - tx: left_tx, - } = left; - let Self { - rx: mut right_rx, - tx: right_tx, - } = right; - - let left_to_right = async move { - while let Some(frame) = left_rx.next().await { - frame.inspect_messages(&mut left_to_right)?; - right_tx - .unbounded_send(frame) - .map_err(crate::util::internal_error)?; - } - Ok::<(), crate::Error>(()) - }; - let right_to_left = async move { - while let Some(frame) = right_rx.next().await { - frame.inspect_messages(&mut right_to_left)?; - left_tx - .unbounded_send(frame) - .map_err(crate::util::internal_error)?; - } - Ok::<(), crate::Error>(()) + fn charged_for(&self, budget: &Arc) -> usize { + let own = if Arc::ptr_eq(&self.inner.budget, budget) { + self.inner.bytes.load(Ordering::Acquire) + } else { + 0 }; - - futures::try_join!(left_to_right, right_to_left)?; - Ok(()) + own + self + .additional + .iter() + .map(|permit| permit.charged_for(budget)) + .sum::() } -} -impl ConnectTo for Channel { - async fn connect_to(self, client: impl ConnectTo) -> Result<(), crate::Error> { - let (client_channel, client_future) = client.into_channel_and_future(); + fn cover_budget( + &self, + budget: &Arc, + bytes: usize, + data: bool, + ) -> Result<(), crate::Error> { + if Arc::ptr_eq(&self.inner.budget, budget) { + self.cover_frame(bytes, data) + } else { + self.additional + .iter() + .find(|permit| permit.charged_for(budget) > 0) + .expect("destination charge exists") + .cover_budget(budget, bytes, data) + } + } - let ((), (), ()) = futures::try_join!( - Channel { - rx: client_channel.rx, - tx: self.tx, - } - .copy(), - Channel { - rx: self.rx, - tx: client_channel.tx, + fn cover_frame(&self, bytes: usize, data: bool) -> Result<(), crate::Error> { + let budget = &self.inner.budget; + // Aggregated batch permits can exceed the maximum size of one frame. + // That must not allow an oversized frame to bypass the per-frame limit. + if bytes > budget.limits.max_frame_bytes { + return Err( + crate::Error::invalid_request().data("frame exceeds connection byte budget") + ); + } + if data && !self.inner.data { + return Err(crate::Error::invalid_request() + .data("data frame cannot grow a control reservation")); + } + let charged = self.charged_for(budget); + if bytes <= charged { + return Ok(()); + } + let delta = bytes - charged; + let mut state = budget.state.lock().expect("frame budget poisoned"); + if state + .used + .checked_add(delta) + .is_none_or(|used| used > budget.limits.max_queued_bytes) + || data + && state.data_used.checked_add(delta).is_none_or(|used| { + used > budget + .limits + .max_queued_bytes + .saturating_sub(budget.limits.max_frame_bytes) + }) + { + return Err(crate::Error::invalid_request() + .data("outgoing frame exceeds admitted byte capacity")); + } + state.used += delta; + if self.inner.data { + state.data_used += delta; + } + self.inner.bytes.fetch_add(delta, Ordering::Release); + Ok(()) + } +} + +#[derive(Debug)] +struct FramePermitInner { + budget: Arc, + bytes: std::sync::atomic::AtomicUsize, + data: bool, +} + +impl Drop for FramePermitInner { + fn drop(&mut self) { + let mut state = self.budget.state.lock().expect("frame budget poisoned"); + let bytes = self.bytes.load(Ordering::Acquire); + state.used -= bytes; + if self.data { + state.data_used -= bytes; + } + let waiters = state + .waiters + .iter() + .map(|(_, waker)| waker.clone()) + .collect::>(); + drop(state); + for waker in waiters { + waker.wake(); + } + } +} + +#[derive(Debug)] +struct FrameBudget { + limits: ConnectionLimits, + state: Mutex, +} + +#[derive(Debug, Default)] +struct FrameBudgetState { + used: usize, + data_used: usize, + waiters: Vec<(usize, Waker)>, + next_waiter: usize, +} + +struct FrameWaiter { + budget: Arc, + id: Option, +} + +impl Drop for FrameWaiter { + fn drop(&mut self) { + if let Some(id) = self.id { + self.budget + .state + .lock() + .expect("frame budget poisoned") + .waiters + .retain(|(registered, _)| *registered != id); + } + } +} + +impl FrameBudget { + fn try_reserve(self: &Arc, bytes: usize, data: bool) -> Option { + let mut state = self.state.lock().expect("frame budget poisoned"); + if bytes > self.limits.max_frame_bytes + || state.used.checked_add(bytes)? > self.limits.max_queued_bytes + || data + && state.data_used.checked_add(bytes)? + > self + .limits + .max_queued_bytes + .saturating_sub(self.limits.max_frame_bytes) + { + return None; + } + state.used += bytes; + if data { + state.data_used += bytes; + } + Some(FramePermit { + inner: Arc::new(FramePermitInner { + budget: self.clone(), + bytes: std::sync::atomic::AtomicUsize::new(bytes), + data, + }), + additional: Vec::new(), + }) + } + + async fn reserve( + self: &Arc, + bytes: usize, + data: bool, + ) -> Result { + if bytes > self.limits.max_frame_bytes || bytes > self.limits.max_queued_bytes { + return Err( + crate::Error::invalid_request().data("frame exceeds connection byte budget") + ); + } + if data + && bytes + > self + .limits + .max_queued_bytes + .saturating_sub(self.limits.max_frame_bytes) + { + return Err( + crate::Error::invalid_request().data("data frame exceeds connection byte budget") + ); + } + let mut waiter = FrameWaiter { + budget: self.clone(), + id: None, + }; + future::poll_fn(|cx| { + let mut state = self.state.lock().expect("frame budget poisoned"); + if state + .used + .checked_add(bytes) + .is_some_and(|used| used <= self.limits.max_queued_bytes) + && (!data + || state.data_used.checked_add(bytes).is_some_and(|used| { + used <= self + .limits + .max_queued_bytes + .saturating_sub(self.limits.max_frame_bytes) + })) + { + state.used += bytes; + if data { + state.data_used += bytes; + } + Poll::Ready(Ok(FramePermit { + inner: Arc::new(FramePermitInner { + budget: self.clone(), + bytes: std::sync::atomic::AtomicUsize::new(bytes), + data, + }), + additional: Vec::new(), + })) + } else { + if let Some(id) = waiter.id { + let (_, waker) = state + .waiters + .iter_mut() + .find(|(registered, _)| *registered == id) + .expect("registered budget waiter"); + waker.clone_from(cx.waker()); + } else { + let id = state.next_waiter; + state.next_waiter = state.next_waiter.wrapping_add(1); + state.waiters.push((id, cx.waker().clone())); + waiter.id = Some(id); + } + Poll::Pending + } + }) + .await + } +} + +/// A frame with its retained byte admission. Forward this envelope rather +/// than extracting the frame when placing data into another queue. +#[derive(Debug)] +pub struct BudgetedFrame { + frame: TransportFrame, + permit: FramePermit, +} + +impl BudgetedFrame { + /// Borrow the frame without releasing admission. + #[must_use] + pub fn frame(&self) -> &TransportFrame { + &self.frame + } + + /// Borrow the charge when retaining metadata derived from this frame. + /// Cloning the permit retains admission without cloning the payload. + #[must_use] + pub fn permit(&self) -> &FramePermit { + &self.permit + } + + /// Separate the frame and permit for deferred processing. Keep the permit + /// alongside any deferred output until that output has been consumed. + #[must_use] + pub fn into_parts(self) -> (TransportFrame, FramePermit) { + (self.frame, self.permit) + } + + /// Release the frame's admission explicitly after consuming it. + #[must_use] + pub fn into_frame(self) -> TransportFrame { + self.frame + } +} + +/// Pollable receive half of an in-memory duplex. +#[derive(Debug)] +pub struct FrameReceiver(std::pin::Pin>>); + +impl futures::Stream for FrameReceiver { + type Item = BudgetedFrame; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll> { + self.0.as_mut().poll_next(cx) + } +} + +/// Backpressured frame sink. A synchronous send fails when its finite queue is full; +/// asynchronous producers should use [`SinkExt::send`] instead. +pub struct FrameSender { + tx: async_channel::Sender, + budget: Arc, + pending: Mutex>>>, +} + +impl std::fmt::Debug for FrameSender { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("FrameSender") + .field("tx", &self.tx) + .field("budget", &self.budget) + .finish_non_exhaustive() + } +} + +/// Shared byte admission independent of a channel's send half. Used when an +/// adapter stages frames before forwarding them to the channel sink. +#[derive(Clone, Debug)] +pub struct FrameAdmission(Arc); + +impl FrameAdmission { + /// Limits shared by both halves of the duplex connection. + #[must_use] + pub fn limits(&self) -> ConnectionLimits { + self.0.limits + } + + fn try_reserve_bytes(&self, bytes: usize, data: bool) -> Option { + self.0.try_reserve(bytes, data) + } + + async fn reserve_bytes(&self, bytes: usize, data: bool) -> Result { + self.0.reserve(bytes, data).await + } + + /// Admit a frame before placing it into any staging queue. + pub fn try_admit(&self, frame: TransportFrame) -> Result { + let bytes = frame + .to_json() + .map_err(|_| FrameSendError { + frame: Box::new(frame.clone()), + reason: "cannot serialize outgoing JSON-RPC frame", + })? + .len(); + let Some(permit) = self.0.try_reserve(bytes, !frame.is_control()) else { + return Err(FrameSendError { + frame: Box::new(frame), + reason: "outgoing frame byte capacity exceeded", + }); + }; + Ok(BudgetedFrame { frame, permit }) + } + + /// Wait for byte capacity when staging a frame outside inline dispatch. + pub async fn admit(&self, frame: TransportFrame) -> Result { + let bytes = frame.to_json()?.len(); + let permit = self.0.reserve(bytes, !frame.is_control()).await?; + Ok(BudgetedFrame { frame, permit }) + } +} + +impl Clone for FrameSender { + fn clone(&self) -> Self { + Self { + tx: self.tx.clone(), + budget: self.budget.clone(), + pending: Mutex::new(None), + } + } +} + +/// Failure to admit a frame (capacity, size, or closed receiver). The original +/// frame remains available to callers; nothing is silently discarded. +#[derive(Debug)] +pub struct FrameSendError { + frame: Box, + reason: &'static str, +} + +impl std::fmt::Display for FrameSendError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(self.reason) + } +} + +impl std::error::Error for FrameSendError {} + +impl FrameSendError { + /// Recover the frame that was not admitted. + pub fn into_inner(self) -> TransportFrame { + *self.frame + } +} + +impl FrameSender { + /// Obtain the byte admission handle without retaining this channel sender. + pub fn admission(&self) -> FrameAdmission { + FrameAdmission(self.budget.clone()) + } + + /// Fail immediately rather than blocking a protocol dispatcher on its own output. + pub fn try_send(&self, frame: TransportFrame) -> Result<(), FrameSendError> { + let budgeted = self.admission().try_admit(frame)?; + self.tx.try_send(budgeted).map_err(|error| FrameSendError { + frame: Box::new(error.into_inner().frame), + reason: "outgoing frame queue full or closed", + }) + } + + /// Await byte and frame capacity outside ordered dispatch. + pub async fn send_frame(&self, frame: TransportFrame) -> Result<(), crate::Error> { + let bytes = frame.to_json()?.len(); + let permit = match future::select( + Box::pin(self.budget.reserve(bytes, !frame.is_control())), + Box::pin(self.tx.closed()), + ) + .await + { + Either::Left((result, _)) => result?, + Either::Right(((), _)) => { + return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); + } + }; + self.tx + .send(BudgetedFrame { frame, permit }) + .await + .map_err(crate::util::internal_error) + } + + /// Transfer an application message's charge into its framed representation. + /// Any transform expansion must grow that lease immediately, rather than + /// awaiting capacity held by this very message. + async fn send_admitted( + &self, + frame: TransportFrame, + permit: FramePermit, + ) -> Result<(), crate::Error> { + let bytes = frame.to_json()?.len(); + permit.cover_frame(bytes, !frame.is_control())?; + self.tx + .send(BudgetedFrame { frame, permit }) + .await + .map_err(crate::util::internal_error) + } + + /// Stop accepting frames on this queue. + pub fn close_channel(&self) { + self.tx.close(); + self.pending.lock().expect("frame sender poisoned").take(); + } + + /// Return whether the receiving endpoint has closed. + pub fn is_closed(&self) -> bool { + self.tx.is_closed() + } +} + +impl Sink for FrameSender { + type Error = crate::Error; + + fn poll_ready( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.poll_flush(cx) + } + + fn start_send( + self: std::pin::Pin<&mut Self>, + mut item: BudgetedFrame, + ) -> Result<(), Self::Error> { + let this = self.get_mut(); + if this.tx.is_closed() { + return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); + } + let bytes = item.frame.to_json()?.len(); + let data = !item.frame.is_control(); + if bytes > this.budget.limits.max_frame_bytes { + return Err( + crate::Error::invalid_request().data("frame exceeds connection byte budget") + ); + } + if item.permit.charged_for(&this.budget) > 0 { + item.permit.cover_budget(&this.budget, bytes, data)?; + } else { + let permit = this.budget.try_reserve(bytes, data).ok_or_else(|| { + crate::Error::invalid_request().data("outgoing frame byte capacity exceeded") + })?; + item.permit.join(permit); + } + let mut pending = this.pending.lock().expect("frame sender poisoned"); + if pending.is_some() { + return Err(crate::Error::invalid_request().data("frame sender not ready")); + } + let tx = this.tx.clone(); + *pending = Some(Box::pin(async move { + tx.send(item).await.map_err(crate::util::internal_error) + })); + Ok(()) + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + let this = self.get_mut(); + let mut pending = this.pending.lock().expect("frame sender poisoned"); + let Some(send) = pending.as_mut() else { + return Poll::Ready(if this.tx.is_closed() { + Err(crate::Error::invalid_request().data("outgoing frame queue closed")) + } else { + Ok(()) + }); + }; + match send.as_mut().poll(cx) { + Poll::Pending => Poll::Pending, + Poll::Ready(result) => { + pending.take(); + Poll::Ready(result) + } + } + } + + fn poll_close( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + let this = self.get_mut(); + match std::pin::Pin::new(&mut *this).poll_flush(cx) { + Poll::Ready(Ok(())) => { + this.tx.close(); + Poll::Ready(Ok(())) + } + other => other, + } + } +} + +impl Channel { + /// Create a pair of connected channel endpoints. + /// + /// Frames sent through either endpoint are received by the other endpoint. + #[must_use] + pub fn duplex() -> (Self, Self) { + Self::duplex_with_limits(ConnectionLimits::default()) + } + + /// Create a connected pair sharing one finite byte budget. + #[must_use] + pub fn duplex_with_limits(limits: ConnectionLimits) -> (Self, Self) { + let budget = Arc::new(FrameBudget { + limits, + state: Mutex::new(FrameBudgetState::default()), + }); + let (a_tx, b_rx) = async_channel::bounded(limits.max_queued_frames.max(1)); + let (b_tx, a_rx) = async_channel::bounded(limits.max_queued_frames.max(1)); + ( + Self { + rx: FrameReceiver(Box::pin(a_rx)), + tx: FrameSender { + tx: a_tx, + budget: budget.clone(), + pending: Mutex::new(None), + }, + }, + Self { + rx: FrameReceiver(Box::pin(b_rx)), + tx: FrameSender { + tx: b_tx, + budget, + pending: Mutex::new(None), + }, + }, + ) + } + + /// Copy frames from `rx` to `tx` until the input closes. + /// + /// # Errors + /// + /// Returns an error if the receiving endpoint closes before the input. + pub(crate) async fn copy(mut self) -> Result<(), crate::Error> { + while let Some(frame) = self.rx.next().await { + self.tx + .send(frame) + .await + .map_err(crate::util::internal_error)?; + } + Ok(()) + } + + /// Bridge two endpoints while inspecting every valid message. + /// + /// Observers are invoked in source order, including for each valid member of + /// a batch. The original frame is forwarded unchanged after inspection. + /// + /// # Errors + /// + /// Returns an observer error or an error if a destination closes before its + /// source. + pub async fn bridge_with_inspection( + left: Self, + right: Self, + mut left_to_right: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send, + mut right_to_left: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send, + ) -> Result<(), crate::Error> { + let Self { + rx: mut left_rx, + tx: mut left_tx, + } = left; + let Self { + rx: mut right_rx, + tx: mut right_tx, + } = right; + + let left_to_right = async move { + while let Some(frame) = left_rx.next().await { + frame.frame().inspect_messages(&mut left_to_right)?; + right_tx + .send(frame) + .await + .map_err(crate::util::internal_error)?; + } + Ok::<(), crate::Error>(()) + }; + let right_to_left = async move { + while let Some(frame) = right_rx.next().await { + frame.frame().inspect_messages(&mut right_to_left)?; + left_tx + .send(frame) + .await + .map_err(crate::util::internal_error)?; + } + Ok::<(), crate::Error>(()) + }; + + futures::try_join!(left_to_right, right_to_left)?; + Ok(()) + } +} + +impl ConnectTo for Channel { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), crate::Error> { + let (client_channel, client_future) = client.into_channel_and_future(); + let passive = client_future.is_passive(); + + let outbound = Channel { + rx: client_channel.rx, + tx: self.tx, + } + .copy(); + let inbound = Channel { + rx: self.rx, + tx: client_channel.tx, + } + .copy(); + if passive { + // Neither channel owns the remote application. Preserve half-close: + // input EOF must still allow responses to drain the other way. + futures::try_join!(inbound, outbound)?; + return Ok(()); + } + // Poll output while the client is running: its requests may be needed + // to let either peer finish. A raw Channel has a no-op driver, so driver + // completion alone is not a signal to stop forwarding its input. + let local = async move { + futures::try_join!(client_future, outbound)?; + Ok::<(), crate::Error>(()) + }; + match future::select(Box::pin(local), Box::pin(inbound)).await { + Either::Left((result, _inbound)) => { + // The local client has finished and its accepted output drained. + // Do not also wait for a remote sender that can remain alive. + result + } + Either::Right((result, local)) => { + result?; + local.await + } + } + } + + fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { + (self, crate::ConnectionDriver::passive()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + struct SendOneThenFinish; + + impl ConnectTo for SendOneThenFinish { + async fn connect_to( + self, + peer: impl ConnectTo, + ) -> Result<(), crate::Error> { + let (channel, driver) = peer.into_channel_and_future(); + channel + .tx + .send_frame(TransportFrame::Single(RawJsonRpcMessage::notification( + "finished".into(), + serde_json::json!({}), + )?)) + .await + .map_err(crate::util::internal_error)?; + drop(channel); + driver.await + } + } + + #[tokio::test] + async fn channel_connect_finishes_without_remote_eof_after_local_drain() { + let (local, mut remote) = Channel::duplex(); + let connection = tokio::spawn(ConnectTo::::connect_to( + local, + SendOneThenFinish, + )); + let frame = tokio::time::timeout(std::time::Duration::from_secs(2), remote.rx.next()) + .await + .expect("accepted frame should arrive") + .expect("channel open"); + assert!(matches!( + frame.frame(), + TransportFrame::Single(RawJsonRpcMessage::Notification(_)) + )); + tokio::time::timeout(std::time::Duration::from_secs(2), connection) + .await + .expect("local completion must not wait for remote sender") + .expect("connection task") + .expect("connection result"); + // Retaining the remote sender did not block local completion. The + // completed endpoint has now closed its receiving half. + assert!(remote.tx.is_closed()); + } + + #[tokio::test] + async fn frame_permits_survive_dequeue_until_consumed() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("capacity".into(), serde_json::json!({})).unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let (left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes * 2, + max_queued_bytes: bytes * 4, + max_queued_frames: 8, + }); + left.tx.try_send(frame.clone()).unwrap(); + left.tx.try_send(frame.clone()).unwrap(); + let held = right.rx.next().await.unwrap(); + assert_eq!( + held.frame().to_json().unwrap().len(), + held.permit.charged_bytes() + ); + assert!( + left.tx.try_send(frame.clone()).is_err(), + "dequeue must not release byte admission" + ); + drop(held); + left.tx + .try_send(frame) + .expect("dropping the last permit releases capacity"); + } + + fn capacity_frame() -> TransportFrame { + TransportFrame::Single( + RawJsonRpcMessage::notification("capacity".into(), serde_json::json!({})).unwrap(), + ) + } + + #[test] + fn cloned_frame_senders_do_not_expand_queue_capacity() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes * 2, + max_queued_bytes: bytes * 5000, + max_queued_frames: 2, + }); + let clones = (0..3000).map(|_| left.tx.clone()).collect::>(); + clones[0].try_send(frame.clone()).unwrap(); + clones[1].try_send(frame.clone()).unwrap(); + for tx in &clones { + assert!(tx.try_send(frame.clone()).is_err()); + } + drop(right.rx.next().now_or_never().unwrap()); + clones[2999] + .try_send(frame) + .expect("one dequeue restores precisely one slot"); + } + + #[test] + fn task_and_dynamic_queues_remain_bounded_across_clones() { + for name in ["task", "dynamic"] { + let (tx, mut rx) = admission::channel_with_capacity::(2); + let clones = (0..3000).map(|_| tx.clone()).collect::>(); + clones[0].unbounded_send(0).unwrap(); + clones[1].unbounded_send(1).unwrap(); + assert!( + clones + .iter() + .all(|sender| sender.unbounded_send(2).is_err()), + "{name}" + ); + assert_eq!(rx.next().now_or_never().unwrap(), Some(0)); + clones[2999] + .unbounded_send(3) + .expect("dequeue restores one slot"); + } + } + + #[tokio::test] + async fn imported_frames_obey_destination_frame_and_byte_limits() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (source, mut source_peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes * 2, + max_queued_bytes: bytes * 8, + max_queued_frames: 2, + }); + source.tx.try_send(frame.clone()).unwrap(); + let held = source_peer.rx.next().await.unwrap(); + let lease = held.permit.clone(); + let source_budget = source.tx.budget.clone(); + let (mut smaller, _peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes - 1, + max_queued_bytes: bytes * 8, + max_queued_frames: 2, + }); + assert!(smaller.tx.send(held).await.is_err()); + assert_eq!(source_budget.state.lock().unwrap().used, bytes); + drop(lease); + assert_eq!(source_budget.state.lock().unwrap().used, 0); + + source.tx.try_send(frame.clone()).unwrap(); + let held = source_peer.rx.next().await.unwrap(); + let lease = held.permit.clone(); + let (mut smaller, _peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2 - 1, + max_queued_frames: 2, + }); + assert!(smaller.tx.send(held).await.is_err()); + assert_eq!(source_budget.state.lock().unwrap().used, bytes); + drop(lease); + assert_eq!(source_budget.state.lock().unwrap().used, 0); + source + .tx + .try_send(frame) + .expect("rejected import releases source lease"); + } + + #[tokio::test] + async fn same_budget_frame_handoff_does_not_charge_twice() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (mut left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2, + max_queued_frames: 2, + }); + left.tx.try_send(frame).unwrap(); + let held = right.rx.next().await.unwrap(); + right.tx.send(held).await.unwrap(); + let held = left.rx.next().await.unwrap(); + assert_eq!(left.tx.budget.state.lock().unwrap().used, bytes); + drop(held); + assert_eq!(left.tx.budget.state.lock().unwrap().used, 0); + } + + #[tokio::test] + async fn imported_frame_retains_independent_budget_charges_without_recharging_on_return() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2, + max_queued_frames: 2, + }; + let (source, mut source_peer) = Channel::duplex_with_limits(limits); + let (mut destination, mut destination_peer) = Channel::duplex_with_limits(limits); + source.tx.try_send(frame).unwrap(); + destination + .tx + .send(source_peer.rx.next().await.unwrap()) + .await + .unwrap(); + assert_eq!(source.tx.budget.state.lock().unwrap().used, bytes); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, bytes); + let imported = destination_peer.rx.next().await.unwrap(); + destination_peer.tx.send(imported).await.unwrap(); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, bytes); + drop(destination.rx.next().await.unwrap()); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 0); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 0); + } + + #[test] + fn cancelled_byte_waiters_are_unregistered() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, _right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2, + max_queued_frames: 2, + }); + let _held = left.tx.admission().try_admit(frame.clone()).unwrap(); + for _ in 0..1000 { + { + let waiting = left.tx.budget.reserve(bytes, true); + futures::pin_mut!(waiting); + assert!(waiting.as_mut().now_or_never().is_none()); } - .copy(), - client_future, - )?; - Ok(()) + assert!(left.tx.budget.state.lock().unwrap().waiters.is_empty()); + } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { - (self, Box::pin(future::ready(Ok(())))) + #[tokio::test] + async fn closing_receiver_wakes_byte_blocked_sender() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2, + max_queued_frames: 2, + }); + let _held = left.tx.admission().try_admit(frame.clone()).unwrap(); + let mut waiting = Box::pin(left.tx.send_frame(frame)); + assert!(waiting.as_mut().now_or_never().is_none()); + drop(right); + assert!(waiting.await.is_err()); + assert!(left.tx.budget.state.lock().unwrap().waiters.is_empty()); } -} -#[cfg(test)] -mod tests { - use super::*; + #[test] + fn cancelling_queued_async_send_releases_its_byte_charge() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 4, + max_queued_frames: 1, + }); + left.tx.try_send(frame.clone()).unwrap(); + { + let waiting = left.tx.send_frame(frame.clone()); + futures::pin_mut!(waiting); + assert!(waiting.as_mut().now_or_never().is_none()); + assert_eq!(left.tx.budget.state.lock().unwrap().used, bytes * 2); + } + assert_eq!(left.tx.budget.state.lock().unwrap().used, bytes); + drop(right.rx.next().now_or_never().unwrap()); + assert_eq!(left.tx.budget.state.lock().unwrap().used, 0); + } + + #[test] + fn closing_sender_drops_pending_sink_frame_charge() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (mut left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 4, + max_queued_frames: 1, + }); + left.tx.try_send(frame.clone()).unwrap(); + let admitted = left.tx.admission().try_admit(frame).unwrap(); + std::pin::Pin::new(&mut left.tx) + .start_send(admitted) + .unwrap(); + assert!( + std::pin::Pin::new(&mut left.tx) + .poll_flush(&mut Context::from_waker(futures::task::noop_waker_ref())) + .is_pending() + ); + left.tx.close_channel(); + assert_eq!(left.tx.budget.state.lock().unwrap().used, bytes); + drop(right.rx.next().now_or_never().unwrap()); + assert_eq!(left.tx.budget.state.lock().unwrap().used, 0); + } + + fn application_channel( + admission: FrameAdmission, + ) -> ( + outgoing_actor::OutgoingMessageTx, + admission::Receiver, + ) { + admission::budgeted_channel( + admission, + OutgoingMessage::charged_bytes, + OutgoingMessage::with_permit, + OutgoingMessage::is_control, + OutgoingMessage::is_urgent, + ) + } + + #[test] + fn application_byte_wait_is_interrupted_by_receiver_close() { + let message = || OutgoingMessage::Notification { + untyped: UntypedMessage::new("held", serde_json::json!({})).unwrap(), + }; + let charge = message().charged_bytes().unwrap(); + let (channel, _peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: charge, + max_queued_bytes: charge * 2, + max_queued_frames: 2, + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + tx.unbounded_send(message()).unwrap(); + let held = rx.next().now_or_never().unwrap().unwrap(); + let mut blocked = Box::pin(tx.send(message())); + assert!(blocked.as_mut().now_or_never().is_none()); + drop(rx); + assert!( + blocked + .now_or_never() + .expect("closed queue must wake a byte waiter") + .is_err() + ); + // The held payload deliberately outlives closure of its queue. + drop(held); + } + + #[test] + fn unbudgeted_control_queue_reports_its_configured_capacity() { + let (tx, _rx) = admission::channel_with_capacity::(2); + assert_eq!(tx.queue_capacity(), 2); + assert_eq!(tx.clone().queue_capacity(), 2); + } + + #[test] + fn routed_results_have_independent_retained_charges() { + let limits = ConnectionLimits { + max_frame_bytes: 1024, + max_queued_bytes: 2700, + max_queued_frames: 8, + }; + let (channel, _) = Channel::duplex_with_limits(limits); + let admission = channel.tx.admission(); + let initial = admission.try_reserve_bytes(100, true).unwrap(); + let value = serde_json::json!("x".repeat(700)); + let mut receivers = Vec::new(); + for i in 0..2 { + let (sender, receiver) = oneshot::channel(); + let id = RequestId::Str(format!("response-{i}")); + let pending = PendingReply { + method: "test".into(), + metadata_bytes: None, + role_id: crate::role::UntypedRole.role_id(), + sender, + cancellation_disarm: SentRequestCancellationDisarm::new(), + ordering: ResponseOrdering::default(), + response_route_hook: None, + }; + let (dispatch, _) = incoming_actor::dispatch_from_response( + id, + pending, + Ok(value.clone()), + Some(initial.clone()), + ); + let Dispatch::Response(result, router) = dispatch else { + panic!("response expected") + }; + router.route_with_result(result).unwrap(); + receivers.push(receiver); + } + drop(initial); + let used = admission.0.state.lock().unwrap().used; + assert_eq!(used, 2 * serde_json::to_vec(&value).unwrap().len()); + let first = futures::executor::block_on(receivers.remove(0)).unwrap(); + assert!(first.result.is_ok()); + assert!(admission.0.state.lock().unwrap().used > 0); + drop(first); + drop(receivers); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + + #[test] + fn oversized_transformed_result_fails_without_waiting_on_its_frame() { + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 1000, + max_queued_frames: 8, + }); + let admission = channel.tx.admission(); + let frame = admission.try_reserve_bytes(200, true).unwrap(); + let (sender, receiver) = oneshot::channel(); + let pending = PendingReply { + method: "transform".into(), + metadata_bytes: None, + role_id: crate::role::UntypedRole.role_id(), + sender, + cancellation_disarm: SentRequestCancellationDisarm::new(), + ordering: ResponseOrdering::default(), + response_route_hook: None, + }; + let (dispatch, _) = incoming_actor::dispatch_from_response( + RequestId::Str("transform".into()), + pending, + Ok(serde_json::json!(null)), + Some(frame.clone()), + ); + let Dispatch::Response(_, router) = dispatch else { + panic!("response expected") + }; + router.route(serde_json::json!("x".repeat(400))).unwrap(); + drop(frame); + let received = futures::executor::block_on(receiver).unwrap(); + assert!(received.result.is_err()); + drop(received); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + + #[test] + fn callback_keeps_result_admitted_until_callback_finishes() { + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 1024, + max_queued_bytes: 2700, + max_queued_frames: 8, + }); + let admission = channel.tx.admission(); + let (message_tx, _message_rx) = application_channel(admission.clone()); + let (task_tx, mut task_rx) = admission::channel(); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); + let pending = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx, + task_tx, + dynamic_handler_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending.registrar(), + ProtocolMode::disabled(), + ); + let sent = connection.send_request_to( + crate::role::UntypedRole, + UntypedMessage::new("callback", serde_json::json!({})).unwrap(), + ); + let id = sent.id().clone(); + let frame = admission.try_reserve_bytes(100, true).unwrap(); + let pending_reply = pending.remove(&id).unwrap(); + let (dispatch, _) = incoming_actor::dispatch_from_response( + id, + pending_reply, + Ok(serde_json::json!("x".repeat(500))), + Some(frame.clone()), + ); + let Dispatch::Response(result, router) = dispatch else { + panic!("response expected") + }; + router.route_with_result(result).unwrap(); + drop(frame); + let (finish_tx, finish_rx) = oneshot::channel::<()>(); + sent.on_receiving_result(move |result| async move { + assert!(result.is_ok()); + finish_rx.await.unwrap(); + Ok(()) + }) + .unwrap(); + let task = futures::FutureExt::now_or_never(futures::StreamExt::next(&mut task_rx)) + .unwrap() + .unwrap(); + let mut running = Box::pin(task.run_for_test()); + assert!(running.as_mut().now_or_never().is_none()); + assert!(admission.0.state.lock().unwrap().used >= 500); + finish_tx.send(()).unwrap(); + futures::executor::block_on(running).unwrap(); + // The outgoing frame is still queued; only the callback's result + // charge has been released. + assert!(admission.0.state.lock().unwrap().used < 500); + } + + #[test] + fn cloned_application_senders_respect_item_capacity_independently_of_bytes() { + let message = || OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({})).unwrap(), + }; + let charge = message().charged_bytes().unwrap(); + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: charge * 2, + max_queued_bytes: charge * 6000, + max_queued_frames: 2, + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + let clones = (0..3000).map(|_| tx.clone()).collect::>(); + clones[0].unbounded_send(message()).unwrap(); + clones[1].unbounded_send(message()).unwrap(); + assert!( + clones + .iter() + .all(|sender| sender.unbounded_send(message()).is_err()) + ); + drop(rx.next().now_or_never().unwrap()); + clones[2999].unbounded_send(message()).unwrap(); + } + + #[test] + fn application_payload_is_charged_after_dequeue_until_dropped() { + let message = OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({"value": "123"})).unwrap(), + }; + let charge = message.charged_bytes().unwrap(); + let frame_bytes = charge + 32; + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: frame_bytes, + max_queued_bytes: frame_bytes + 2 * charge, + max_queued_frames: 3, + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({"value": "123"})).unwrap(), + }) + .unwrap(); + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({"value": "123"})).unwrap(), + }) + .unwrap(); + let held = rx.next().now_or_never().unwrap().unwrap(); + assert!( + tx.unbounded_send(message).is_err(), + "dequeue must retain application admission" + ); + drop(held); + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({"value": "123"})).unwrap(), + }) + .expect("capacity is recovered after the retained application message is dropped"); + } + + #[tokio::test] + async fn application_lease_moves_into_writer_frame_without_recharging() { + let message = OutgoingMessage::Notification { + untyped: UntypedMessage::new("handoff", serde_json::json!({"value": "abc"})).unwrap(), + }; + let charge = message.charged_bytes().unwrap(); + let frame_bytes = charge + 16; + let (sender, mut receiver) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: frame_bytes, + max_queued_bytes: frame_bytes + charge, + max_queued_frames: 2, + }); + let (tx, mut application_rx) = application_channel(sender.tx.admission()); + tx.unbounded_send(message).unwrap(); + let OutgoingMessage::Admitted { message, permit } = application_rx.next().await.unwrap() + else { + panic!("application admission must wrap the queued payload"); + }; + let OutgoingMessage::Notification { untyped } = *message else { + panic!("expected notification"); + }; + let frame = TransportFrame::Single(untyped.into_raw_jsonrpc_message(None).unwrap()); + sender.tx.send_admitted(frame, permit).await.unwrap(); + let held = receiver.rx.next().await.unwrap(); + assert!( + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("handoff", serde_json::json!({"value": "abc"})) + .unwrap(), + }) + .is_err(), + "writer-held frame keeps its application charge" + ); + drop(held); + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("handoff", serde_json::json!({"value": "abc"})).unwrap(), + }) + .expect("capacity is recovered after the writer frame is released"); + } + + #[test] + fn cancellation_lane_is_ready_when_data_queue_is_full() { + let ordinary = OutgoingMessage::Notification { + untyped: UntypedMessage::new("ordinary", serde_json::json!({})).unwrap(), + }; + let cancel = OutgoingMessage::Notification { + untyped: UntypedMessage::new( + "$/cancel_request", + serde_json::json!({"requestId":"one"}), + ) + .unwrap(), + }; + let frame_bytes = ordinary + .charged_bytes() + .unwrap() + .max(cancel.charged_bytes().unwrap()) + + 16; + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: frame_bytes, + max_queued_bytes: frame_bytes * 2, + max_queued_frames: 1, + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + tx.unbounded_send(ordinary).unwrap(); + tx.unbounded_send(cancel) + .expect("cancellation has a separate control lane"); + assert!(rx.next().now_or_never().unwrap().unwrap().is_urgent()); + } + + #[test] + fn cancellation_passes_waiting_request_and_saturated_data_lane() { + let (transport, mut peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 1024, + max_queued_bytes: 4096, + max_queued_frames: 1, + }); + let (message_tx, message_rx) = application_channel(transport.tx.admission()); + let (task_tx, _task_rx) = admission::channel(); + let (dynamic_tx, _dynamic_rx) = admission::channel(); + let pending_replies = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx.clone(), + task_tx, + dynamic_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending_replies.registrar(), + ProtocolMode::disabled(), + ); + let (ready_tx, ready_rx) = oneshot::channel(); + let sent = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("not-ready", serde_json::json!({})).unwrap(), + async move { ready_rx.await.map_err(crate::Error::into_internal_error) }, + ); + let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( + message_rx, + pending_replies, + transport.tx, + ProtocolCompat::new(ProtocolMode::disabled()), + IncomingClosed::new(), + )); + assert!(actor.as_mut().now_or_never().is_none()); + message_tx + .unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("data", serde_json::json!({})).unwrap(), + }) + .unwrap(); + connection + .send_cancel_request(sent.id().clone()) + .expect("urgent lane should remain available"); + assert!(actor.as_mut().now_or_never().is_none()); + let error = sent + .block_task() + .now_or_never() + .expect("cancel settles without readiness") + .expect_err("unpublished request must be cancelled locally"); + assert_eq!(error.code, crate::ErrorCode::RequestCancelled); + assert!( + ready_tx.send(()).is_err(), + "cancelled readiness future must be dropped" + ); + let data = peer.rx.next().now_or_never().unwrap().unwrap(); + assert!( + matches!(data.frame(), TransportFrame::Single(RawJsonRpcMessage::Notification(n)) + if n.method.as_ref() == "data") + ); + assert!( + peer.rx.next().now_or_never().is_none(), + "never publish a request after its cancellation" + ); + } + + #[test] + fn cancellation_of_queued_request_does_not_wait_for_unrelated_readiness() { + let (transport, mut peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 1024, + max_queued_bytes: 8192, + max_queued_frames: 1, + }); + let (message_tx, message_rx) = application_channel(transport.tx.admission()); + let (task_tx, _task_rx) = admission::channel(); + let (dynamic_tx, _dynamic_rx) = admission::channel(); + let pending_replies = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx, + task_tx, + dynamic_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending_replies.registrar(), + ProtocolMode::disabled(), + ); + let (ready_tx, ready_rx) = oneshot::channel(); + let first = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("first", serde_json::json!({})).unwrap(), + async move { ready_rx.await.map_err(crate::Error::into_internal_error) }, + ); + let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( + message_rx, + pending_replies, + transport.tx, + ProtocolCompat::new(ProtocolMode::disabled()), + IncomingClosed::new(), + )); + assert!(actor.as_mut().now_or_never().is_none()); + let second = connection.send_request_to( + crate::role::UntypedRole, + UntypedMessage::new("second", serde_json::json!({})).unwrap(), + ); + second.cancel().unwrap(); + assert!(actor.as_mut().now_or_never().is_none()); + let error = second + .block_task() + .now_or_never() + .expect("queued cancellation cannot wait for first") + .expect_err("second request was never published"); + assert_eq!(error.code, crate::ErrorCode::RequestCancelled); + assert!(peer.rx.next().now_or_never().is_none()); + + ready_tx.send(()).unwrap(); + assert!(actor.as_mut().now_or_never().is_none()); + let frame = peer.rx.next().now_or_never().unwrap().unwrap(); + assert!( + matches!(frame.frame(), TransportFrame::Single(RawJsonRpcMessage::Request(r)) + if r.method.as_ref() == "first") + ); + drop(frame); + assert!( + peer.rx.next().now_or_never().is_none(), + "second must not run later" + ); + first.detach(); + } #[cfg(feature = "unstable_protocol_v2")] fn connection_with_task_receiver() -> ( ConnectionTo, - mpsc::UnboundedReceiver, + admission::SimpleReceiver, ) { - let (message_tx, _message_rx) = mpsc::unbounded(); - let (task_tx, task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, _message_rx) = admission::channel(); + let (task_tx, task_rx) = admission::channel(); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -6708,9 +8560,9 @@ mod tests { #[cfg(feature = "unstable_protocol_v2")] #[test] fn v2_proxy_rejects_explicitly_prewrapped_initialize_request() { - let (message_tx, message_rx) = mpsc::unbounded(); - let (task_tx, _task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, message_rx) = admission::channel(); + let (task_tx, _task_rx) = admission::channel(); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -6734,12 +8586,21 @@ mod tests { }; let sent = connection.send_request_to(Agent, request); - let (transport_tx, mut transport_rx) = mpsc::unbounded(); + let ( + Channel { + tx: transport_tx, .. + }, + Channel { + rx: mut transport_rx, + .. + }, + ) = Channel::duplex(); let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( message_rx, pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::v2_proxy()), + IncomingClosed::new(), )); assert!( actor.as_mut().now_or_never().is_none(), @@ -6859,11 +8720,11 @@ mod tests { fn connection_with_dynamic_handler_receiver() -> ( ConnectionTo, - mpsc::UnboundedReceiver>, + admission::SimpleReceiver>, ) { - let (message_tx, _message_rx) = mpsc::unbounded(); - let (task_tx, _task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, _message_rx) = admission::channel(); + let (task_tx, _task_rx) = admission::channel(); + let (dynamic_handler_tx, dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -6900,12 +8761,12 @@ mod tests { fn connection_for_response_hook_tests() -> ( ConnectionTo, - mpsc::UnboundedReceiver, + admission::SimpleReceiver, PendingReplies, ) { - let (message_tx, message_rx) = mpsc::unbounded(); - let (task_tx, _task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, message_rx) = admission::channel(); + let (task_tx, _task_rx) = admission::channel(); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -6925,6 +8786,138 @@ mod tests { ) } + fn budgeted_request_connection( + limits: ConnectionLimits, + ) -> ( + ConnectionTo, + admission::Receiver, + PendingReplies, + FrameAdmission, + ) { + let (channel, _) = Channel::duplex_with_limits(limits); + let admission = channel.tx.admission(); + let (message_tx, message_rx) = application_channel(admission.clone()); + let (task_tx, _task_rx) = admission::channel(); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); + let pending = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx, + task_tx, + dynamic_handler_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending.registrar(), + ProtocolMode::disabled(), + ); + (connection, message_rx, pending, admission) + } + + #[test] + fn pending_request_metadata_remains_charged_after_queue_consumption() { + let (connection, mut rx, pending, admission) = + budgeted_request_connection(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 1312, + max_queued_frames: 32, + }); + let method = "m".repeat(140); + let request = || UntypedMessage::new(&method, serde_json::json!({})).unwrap(); + let first = connection.send_request_to(crate::role::UntypedRole, request()); + let first_id = first.id().clone(); + let queued = futures::FutureExt::now_or_never(futures::StreamExt::next(&mut rx)) + .unwrap() + .unwrap(); + drop(queued); + assert!(admission.0.state.lock().unwrap().used > 0); + let second = connection.send_request_to(crate::role::UntypedRole, request()); + assert!(futures::executor::block_on(second.block_task()).is_err()); + assert!(pending.remove(&first_id).is_some()); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + drop(first); + } + + #[test] + fn rejected_admitted_request_fails_without_leaking_pending_reply() { + let (connection, mut rx, pending, admission) = + budgeted_request_connection(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 2048, + max_queued_frames: 1, + }); + admission::ReceiverClose::close(&mut rx); + let sent = connection.send_request_to( + crate::role::UntypedRole, + UntypedMessage::new("rejected", serde_json::json!({})).unwrap(), + ); + assert!(!pending.contains(sent.id())); + assert!(futures::executor::block_on(sent.block_task()).is_err()); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + + #[test] + fn cancelling_request_releases_pending_metadata_after_queue_consumption() { + let (connection, mut rx, pending, admission) = + budgeted_request_connection(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 1312, + max_queued_frames: 8, + }); + let sent = connection.send_request_to( + crate::role::UntypedRole, + UntypedMessage::new(&"m".repeat(140), serde_json::json!({})).unwrap(), + ); + let id = sent.id().clone(); + drop(futures::FutureExt::now_or_never(futures::StreamExt::next(&mut rx)).unwrap()); + assert!(pending.contains(&id)); + drop(sent); + assert!(!pending.contains(&id)); + // Drop the cancellation notification too; no payload remains admitted. + drop(futures::FutureExt::now_or_never(futures::StreamExt::next(&mut rx)).unwrap()); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + + #[test] + fn incoming_eof_releases_many_pending_method_charges_after_error_consumption() { + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 2100, + max_queued_frames: 32, + }); + let admission = channel.tx.admission(); + let pending = PendingReplies::with_capacity(32); + let mut receivers = Vec::new(); + for i in 0..7 { + let method = "m".repeat(120); + let id = RequestId::Str(format!("{i:036}")); + let charge = admission + .try_reserve_bytes(method.len() + 36 + 64, true) + .unwrap(); + let (sender, receiver) = oneshot::channel(); + assert!(pending.registrar().subscribe( + id, + PendingReply { + method, + metadata_bytes: Some(charge), + role_id: crate::role::UntypedRole.role_id(), + sender, + cancellation_disarm: SentRequestCancellationDisarm::new(), + ordering: ResponseOrdering::default(), + response_route_hook: None, + }, + &IncomingClosed::new(), + )); + receivers.push(receiver); + } + assert!(admission.try_reserve_bytes(220, true).is_none()); + assert_eq!(pending.close_incoming(), 7); + assert!( + admission.0.state.lock().unwrap().used > 0, + "failed results still own their method text" + ); + drop(receivers); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + #[cfg(feature = "unstable_protocol_v2")] fn route_test_response( request_id: RequestId, @@ -6935,7 +8928,7 @@ mod tests { .remove(&request_id) .expect("the request should have a pending reply"); let (dispatch, _) = - incoming_actor::dispatch_from_response(request_id, pending_reply, result); + incoming_actor::dispatch_from_response(request_id, pending_reply, result, None); let Dispatch::Response(result, router) = dispatch else { panic!("expected a response dispatch"); }; @@ -7068,12 +9061,21 @@ mod tests { async move { ready_rx.await.map_err(crate::Error::into_internal_error) }, ); - let (transport_tx, mut transport_rx) = mpsc::unbounded(); + let ( + Channel { + tx: transport_tx, .. + }, + Channel { + rx: mut transport_rx, + .. + }, + ) = Channel::duplex(); let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( message_rx, pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::disabled()), + IncomingClosed::new(), )); assert!( @@ -7098,13 +9100,43 @@ mod tests { .expect("the ready request should be published") .expect("the transport queue should remain open"); assert!(matches!( - frame, + frame.frame(), TransportFrame::Single(RawJsonRpcMessage::Request(_)) )); drop(sent); } + #[test] + fn pending_outgoing_readiness_is_cancelled_on_shutdown() { + let (connection, message_rx, pending_replies) = connection_for_response_hook_tests(); + let sent = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("waiting", serde_json::json!({})).unwrap(), + future::pending::>(), + ); + let ( + Channel { + tx, + rx: mut transport_rx, + }, + _peer, + ) = Channel::duplex(); + let shutdown = IncomingClosed::new(); + let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( + message_rx, + pending_replies, + tx, + ProtocolCompat::new(ProtocolMode::disabled()), + shutdown.clone(), + )); + assert!(actor.as_mut().now_or_never().is_none()); + shutdown.begin_close(); + assert!(actor.as_mut().now_or_never().is_none()); + assert!(transport_rx.next().now_or_never().is_none()); + assert!(futures::executor::block_on(sent.block_task()).is_err()); + } + #[test] fn ordered_blocking_transform_precedes_response_acknowledgment() { let (connection, _message_rx, pending_replies) = connection_for_response_hook_tests(); @@ -7121,6 +9153,7 @@ mod tests { request_id, pending_reply, Err(crate::Error::invalid_params()), + None, ); let Dispatch::Response(result, router) = dispatch else { panic!("expected a response dispatch"); @@ -7180,12 +9213,21 @@ mod tests { }, ); - let (transport_tx, mut transport_rx) = mpsc::unbounded(); + let ( + Channel { + tx: transport_tx, .. + }, + Channel { + rx: mut transport_rx, + .. + }, + ) = Channel::duplex(); let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( message_rx, pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::disabled()), + IncomingClosed::new(), )); assert!( @@ -7205,9 +9247,9 @@ mod tests { #[test] fn ordered_request_is_marked_before_entering_outgoing_queue() { - let (message_tx, mut message_rx) = mpsc::unbounded(); - let (task_tx, mut task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, mut message_rx) = admission::channel(); + let (task_tx, mut task_rx) = admission::channel(); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -7250,6 +9292,7 @@ mod tests { request_id, pending_reply, Ok(serde_json::json!({"ok": true})), + None, ); let Dispatch::Response(result, router) = dispatch else { panic!("expected a response dispatch"); @@ -7282,7 +9325,7 @@ mod tests { } fn next_dynamic_handler_message( - receiver: &mut mpsc::UnboundedReceiver>, + receiver: &mut (impl futures::Stream> + Unpin), ) -> Option> { futures::FutureExt::now_or_never(futures::StreamExt::next(receiver)) .expect("dynamic-handler receiver should be ready") @@ -7291,9 +9334,9 @@ mod tests { #[cfg(feature = "unstable_protocol_v2")] #[test] fn v2_dynamic_handler_guard_registers_and_removes_handler() { - let (message_tx, _message_rx) = mpsc::unbounded(); - let (task_tx, _task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, mut dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, _message_rx) = admission::channel(); + let (task_tx, _task_rx) = admission::channel(); + let (dynamic_handler_tx, mut dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); diff --git a/src/agent-client-protocol/src/jsonrpc/admission.rs b/src/agent-client-protocol/src/jsonrpc/admission.rs new file mode 100644 index 00000000..8bb7f1ce --- /dev/null +++ b/src/agent-client-protocol/src/jsonrpc/admission.rs @@ -0,0 +1,274 @@ +//! Finite queues for synchronously invoked dispatcher APIs. +//! +//! Dispatch callbacks cannot await capacity: the receiver may depend on that +//! callback returning. External producers can await `send` instead. +use futures::Stream; +use std::pin::Pin; +use std::sync::Arc; +use std::task::{Context, Poll}; + +use super::{FrameAdmission, FramePermit}; + +pub const QUEUE_CAPACITY: usize = 32; + +pub struct Sender { + inner: Arc>, +} + +struct SenderInner { + tx: async_channel::Sender, + urgent_tx: Option>, + admission: Option>, + capacity: usize, +} + +struct Admission { + budget: FrameAdmission, + measure: fn(&T) -> Result, + attach: fn(T, FramePermit) -> T, + control: fn(&T) -> bool, + urgent: fn(&T) -> bool, +} + +impl Clone for Sender { + fn clone(&self) -> Self { + Self { + inner: self.inner.clone(), + } + } +} + +impl std::fmt::Debug for Sender { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("AdmissionSender").finish_non_exhaustive() + } +} + +impl Sender { + pub fn byte_admission(&self) -> Option { + self.inner + .admission + .as_ref() + .map(|admission| admission.budget.clone()) + } + + pub fn queue_capacity(&self) -> usize { + self.inner.capacity + } + + pub fn unbounded_send(&self, item: T) -> Result<(), SendError> { + // Preserve the order of a queued request followed by its cancellation + // when the ordinary lane has room. Use the bypass lane if ordinary + // capacity is exhausted or its consumer is blocked on readiness. + let urgent = self.inner.admission.as_ref().is_some_and(|admission| { + (admission.urgent)(&item) && (self.inner.tx.is_empty() || self.inner.tx.is_full()) + }); + let item = if let Some(admission) = &self.inner.admission { + let bytes = (admission.measure)(&item).map_err(|error| SendError { + item: None, + reason: error.to_string(), + }); + // Preserve ownership of the rejected message even when sizing fails. + let bytes = match bytes { + Ok(bytes) => bytes, + Err(error) => { + return Err(SendError { + item: Some(item), + ..error + }); + } + }; + let permit = admission + .budget + .try_reserve_bytes(bytes, !(admission.control)(&item)) + .ok_or_else(|| SendError { + item: None, + reason: "outgoing application byte capacity exceeded".into(), + }); + let permit = match permit { + Ok(permit) => permit, + Err(error) => { + return Err(SendError { + item: Some(item), + ..error + }); + } + }; + (admission.attach)(item, permit) + } else { + item + }; + let tx = if urgent { + self.inner.urgent_tx.as_ref().expect("urgent lane exists") + } else { + &self.inner.tx + }; + tx.try_send(item).map_err(|error| SendError { + item: Some(error.into_inner()), + reason: "outgoing application queue full or closed".into(), + }) + } + + pub async fn send(&self, item: T) -> Result<(), crate::Error> { + let urgent = self.inner.admission.as_ref().is_some_and(|admission| { + (admission.urgent)(&item) && (self.inner.tx.is_empty() || self.inner.tx.is_full()) + }); + let tx = if urgent { + self.inner.urgent_tx.as_ref().expect("urgent lane exists") + } else { + &self.inner.tx + }; + let item = if let Some(admission) = &self.inner.admission { + let bytes = (admission.measure)(&item)?; + let reserve = admission + .budget + .reserve_bytes(bytes, !(admission.control)(&item)); + let permit = + match futures::future::select(Box::pin(reserve), Box::pin(tx.closed())).await { + futures::future::Either::Left((permit, _)) => permit?, + futures::future::Either::Right(_) => { + return Err(crate::util::internal_error( + "outgoing application queue closed", + )); + } + }; + (admission.attach)(item, permit) + } else { + item + }; + tx.send(item).await.map_err(crate::util::internal_error) + } +} + +#[derive(Debug)] +pub struct SendError { + item: Option, + reason: String, +} + +impl SendError { + pub fn into_inner(self) -> T { + self.item.expect("send errors retain their rejected item") + } +} + +impl std::fmt::Display for SendError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(&self.reason) + } +} + +impl std::error::Error for SendError {} + +#[cfg(test)] +pub fn channel() -> (Sender, SimpleReceiver) { + channel_with_capacity(QUEUE_CAPACITY) +} + +pub fn channel_with_capacity(capacity: usize) -> (Sender, SimpleReceiver) { + let capacity = capacity.max(1); + let (tx, rx) = async_channel::bounded(capacity); + ( + Sender { + inner: Arc::new(SenderInner { + tx, + urgent_tx: None, + admission: None, + capacity, + }), + }, + SimpleReceiver(Box::pin(rx)), + ) +} + +pub struct SimpleReceiver(Pin>>); + +impl Stream for SimpleReceiver { + type Item = T; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.0.as_mut().poll_next(cx) + } +} + +pub(super) trait ReceiverClose: Stream { + fn close(&mut self); + fn poll_urgent(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Pending + } +} + +impl ReceiverClose for SimpleReceiver { + fn close(&mut self) { + self.0.close(); + } +} + +pub struct Receiver { + normal: SimpleReceiver, + urgent: SimpleReceiver, +} + +impl Stream for Receiver { + type Item = T; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + let urgent_closed = match Pin::new(&mut this.urgent).poll_next(cx) { + Poll::Ready(Some(item)) => return Poll::Ready(Some(item)), + Poll::Ready(None) => true, + Poll::Pending => false, + }; + match Pin::new(&mut this.normal).poll_next(cx) { + Poll::Ready(Some(item)) => Poll::Ready(Some(item)), + Poll::Ready(None) if urgent_closed => Poll::Ready(None), + _ => Poll::Pending, + } + } +} + +impl ReceiverClose for Receiver { + fn close(&mut self) { + self.normal.close(); + self.urgent.close(); + } + + fn poll_urgent(&mut self, cx: &mut Context<'_>) -> Poll> { + match Pin::new(&mut self.urgent).poll_next(cx) { + Poll::Ready(None) => Poll::Pending, + result => result, + } + } +} + +pub fn budgeted_channel( + admission: FrameAdmission, + measure: fn(&T) -> Result, + attach: fn(T, FramePermit) -> T, + control: fn(&T) -> bool, + urgent: fn(&T) -> bool, +) -> (Sender, Receiver) { + let capacity = admission.limits().max_queued_frames.max(1); + let (tx, rx) = async_channel::bounded(capacity); + let (urgent_tx, urgent_rx) = async_channel::bounded(capacity); + ( + Sender { + inner: Arc::new(SenderInner { + tx, + urgent_tx: Some(urgent_tx), + admission: Some(Admission { + budget: admission, + measure, + attach, + control, + urgent, + }), + capacity, + }), + }, + Receiver { + normal: SimpleReceiver(Box::pin(rx)), + urgent: SimpleReceiver(Box::pin(urgent_rx)), + }, + ) +} diff --git a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs index 35383168..a8ff2df3 100644 --- a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs @@ -1,6 +1,5 @@ // Types re-exported from crate root use futures::StreamExt as _; -use futures::channel::mpsc; use futures::stream; use futures_concurrency::stream::StreamExt as _; use rustc_hash::FxHashMap; @@ -24,11 +23,11 @@ use crate::jsonrpc::ResponseDestination; use crate::jsonrpc::ResponseDispatch; use crate::jsonrpc::ResponseRouter; use crate::jsonrpc::TransportBatchEntry; -use crate::jsonrpc::TransportFrame; use crate::jsonrpc::dynamic_handler::DynHandleDispatchFrom; use crate::jsonrpc::dynamic_handler::DynamicHandlerMessage; use crate::jsonrpc::outgoing_actor::send_raw_message; use crate::jsonrpc::protocol_compat::ProtocolCompat; +use crate::jsonrpc::{BudgetedFrame, FramePermit, TransportFrame}; use crate::jsonrpc::{is_response_only_shape, raw_is_response_only_shape}; use crate::role::Role; @@ -59,8 +58,8 @@ impl IncomingHandlers { pub(super) async fn incoming_protocol_actor( counterpart: Counterpart, connection: &ConnectionTo, - transport_rx: mpsc::UnboundedReceiver, - dynamic_handler_rx: mpsc::UnboundedReceiver>, + transport_rx: super::FrameReceiver, + dynamic_handler_rx: super::admission::SimpleReceiver>, pending_replies: PendingReplies, handlers: IncomingHandlers< impl HandleDispatchFrom, @@ -85,7 +84,7 @@ pub(super) async fn incoming_protocol_actor( let mut dynamic_handlers: FxHashMap>> = FxHashMap::default(); - let mut pending_messages: Vec = vec![]; + let mut pending_messages: Vec = vec![]; let request_cancellations = super::RequestCancellationRegistry::new(); let mut on_close = Some(on_close); @@ -128,6 +127,7 @@ pub(super) async fn incoming_protocol_actor( } IncomingProtocolMsg::Transport(frame) => { + let (frame, permit) = frame.into_parts(); let (entries, batch_completion) = frame_entries(frame); for (message, destination) in entries { match message { @@ -156,6 +156,7 @@ pub(super) async fn incoming_protocol_actor( &mut handler, &mut pending_messages, &request_cancellations, + permit.clone(), ) .await?; } @@ -196,6 +197,7 @@ pub(super) async fn incoming_protocol_actor( &mut handler, &mut pending_messages, &request_cancellations, + permit.clone(), ) .await?; } @@ -215,8 +217,12 @@ pub(super) async fn incoming_protocol_actor( if let Some(pending_reply) = pending_replies.remove(&id) { let result = protocol_compat .incoming_response(&pending_reply.method, result); - let (dispatch, response_dispatch) = - dispatch_from_response(id, pending_reply, result); + let (dispatch, response_dispatch) = dispatch_from_response( + id, + pending_reply, + result, + Some(permit.clone()), + ); dispatch_dispatch( counterpart.clone(), connection, @@ -225,6 +231,7 @@ pub(super) async fn incoming_protocol_actor( &mut handler, &mut pending_messages, &request_cancellations, + permit.clone(), ) .await?; if let Some(ack_rx) = response_dispatch.complete() { @@ -260,6 +267,13 @@ pub(super) async fn incoming_protocol_actor( } message @ (IncomingProtocolMsg::Transport(_) | IncomingProtocolMsg::TransportClosed) => { + if queued_transport_messages.len() + >= connection.message_tx.queue_capacity() + { + return Err(crate::util::internal_error( + "transport frames exceed barrier queue capacity", + )); + } queued_transport_messages.push_back(message); } } @@ -303,14 +317,18 @@ async fn handle_dynamic_handler_message( message: DynamicHandlerMessage, connection: &ConnectionTo, dynamic_handlers: &mut FxHashMap>>, - pending_messages: &mut Vec, + pending_messages: &mut Vec, ) -> Result<(), crate::Error> { match message { DynamicHandlerMessage::AddDynamicHandler(uuid, mut handler) => { // Before adding the new handler, give it a chance to process // any pending messages. let mut new_pending_messages = vec![]; - for pending_message in std::mem::take(pending_messages) { + for DeferredDispatch { + dispatch: pending_message, + permit, + } in std::mem::take(pending_messages) + { tracing::trace!(method = pending_message.method(), handler = ?handler.dyn_describe_chain(), "Retrying message"); let reply_target = pending_message.handler_error_target(); let handler_attempt = reply_target @@ -329,7 +347,10 @@ async fn handle_dynamic_handler_message( retry: _, }) => { tracing::trace!(method = m.method(), handler = ?handler.dyn_describe_chain(), "Message not handled"); - new_pending_messages.push(m); + new_pending_messages.push(DeferredDispatch { + dispatch: m, + permit: permit.clone(), + }); } Err(err) => { tracing::warn!(?err, handler = ?handler.dyn_describe_chain(), "Dynamic handler errored on pending message"); @@ -341,6 +362,13 @@ async fn handle_dynamic_handler_message( *pending_messages = new_pending_messages; // Add handler so it will be used for future incoming messages. + if !dynamic_handlers.contains_key(&uuid) + && dynamic_handlers.len() >= connection.message_tx.queue_capacity() + { + return Err(crate::util::internal_error( + "dynamic handler capacity exceeded", + )); + } dynamic_handlers.insert(uuid, handler); } DynamicHandlerMessage::RemoveDynamicHandler(uuid) => { @@ -356,11 +384,16 @@ async fn handle_dynamic_handler_message( #[derive(Debug)] enum IncomingProtocolMsg { - Transport(TransportFrame), + Transport(BudgetedFrame), TransportClosed, DynamicHandler(DynamicHandlerMessage), } +struct DeferredDispatch { + dispatch: Dispatch, + permit: FramePermit, +} + fn frame_entries( frame: TransportFrame, ) -> ( @@ -472,11 +505,17 @@ pub(super) fn dispatch_from_response( id: RequestId, pending_reply: PendingReply, result: Result, + frame_bytes: Option, ) -> (Dispatch, ResponseDispatch) { let response_dispatch = ResponseDispatch::default(); // Create a Dispatch::Response with a ResponseRouter that routes to the oneshot - let router = ResponseRouter::new(id.clone(), pending_reply, response_dispatch.clone()); + let router = ResponseRouter::new( + id.clone(), + pending_reply, + response_dispatch.clone(), + frame_bytes, + ); (Dispatch::Response(result, router), response_dispatch) } @@ -485,14 +524,19 @@ pub(super) fn dispatch_from_response( fields(method = dispatch.method()), level = "trace", )] +#[expect( + clippy::too_many_arguments, + reason = "one dispatch carries its retained frame admission" +)] async fn dispatch_dispatch( counterpart: Counterpart, connection: &ConnectionTo, mut dispatch: Dispatch, dynamic_handlers: &mut FxHashMap>>, handler: &mut impl HandleDispatchFrom, - pending_messages: &mut Vec, + pending_messages: &mut Vec, request_cancellations: &super::RequestCancellationRegistry, + permit: FramePermit, ) -> Result<(), crate::Error> { tracing::trace!(?dispatch, "dispatch_dispatch"); @@ -607,7 +651,15 @@ async fn dispatch_dispatch( ?method, "Retrying message as new dynamic handlers are added" ); - pending_messages.push(dispatch); + if pending_messages.len() >= connection.message_tx.queue_capacity() { + return handle_handler_error( + connection, + error_target, + method, + crate::util::internal_error("pending dispatch capacity exceeded"), + ); + } + pending_messages.push(DeferredDispatch { dispatch, permit }); Ok(()) } else { match dispatch { @@ -617,6 +669,13 @@ async fn dispatch_dispatch( } Dispatch::Request(_, responder) => { tracing::info!(?method, "Rejecting request with error, no handler"); + #[cfg(feature = "unstable_mcp_over_acp")] + if method == "mcp/message" { + return responder.respond_with_error(crate::Error::new( + crate::mcp_server::MCP_SERVER_UNAVAILABLE, + "MCP server unavailable", + )); + } responder.respond_with_error(crate::Error::method_not_found().data(method)) } Dispatch::Response(result, router) => { diff --git a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs index effe164c..f184b847 100644 --- a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs @@ -1,12 +1,15 @@ // Types re-exported from crate root use futures::StreamExt as _; -use futures::channel::mpsc; +use futures::future; +use std::task::Poll; use crate::jsonrpc::protocol_compat::ProtocolCompat; -use crate::jsonrpc::{OutgoingMessage, PendingReplies, RawJsonRpcMessage, TransportFrame}; +use crate::jsonrpc::{ + FramePermit, OutgoingMessage, PendingReplies, RawJsonRpcMessage, TransportFrame, UntypedMessage, +}; use crate::schema::v1::RequestId; -pub type OutgoingMessageTx = mpsc::UnboundedSender; +pub type OutgoingMessageTx = super::admission::Sender; pub(crate) fn send_raw_message( tx: &OutgoingMessageTx, @@ -17,6 +20,45 @@ pub(crate) fn send_raw_message( .map_err(crate::util::internal_error) } +async fn publish( + tx: &super::FrameSender, + frame: TransportFrame, + permit: Option, +) -> Result<(), crate::Error> { + match permit { + Some(permit) => tx.send_admitted(frame, permit).await, + None => tx.send_frame(frame).await, + } + .map_err(crate::Error::into_internal_error) +} + +async fn publish_notification( + tx: &super::FrameSender, + protocol_compat: &ProtocolCompat, + pending_replies: &PendingReplies, + untyped: UntypedMessage, + permit: Option, +) -> Result<(), crate::Error> { + if let Some(id) = super::outgoing_cancellation_id(&untyped) + && pending_replies.cancel_unpublished(&id) + { + return Ok(()); + } + let messages = protocol_compat.outgoing_notification(untyped)?; + // ProtocolCompat currently emits exactly one notification. A future + // expansion needs separately admitted charges for each additional output. + if messages.len() > 1 { + return Err(crate::util::internal_error( + "notification expansion exceeds application admission", + )); + } + if let Some(untyped) = messages.into_iter().next() { + let message = untyped.into_raw_jsonrpc_message(None)?; + publish(tx, TransportFrame::Single(message), permit).await?; + } + Ok(()) +} + /// Outgoing protocol actor: Converts application-level OutgoingMessage to protocol-level RawJsonRpcMessage. /// /// This actor handles JSON-RPC protocol semantics: @@ -25,15 +67,20 @@ pub(crate) fn send_raw_message( /// /// This is the protocol layer - it has no knowledge of how messages are transported. pub(super) async fn outgoing_protocol_actor( - mut outgoing_rx: mpsc::UnboundedReceiver, + mut outgoing_rx: impl Unpin + super::admission::ReceiverClose, pending_replies: PendingReplies, - transport_tx: mpsc::UnboundedSender, + transport_tx: super::FrameSender, protocol_compat: ProtocolCompat, + shutdown: super::IncomingClosed, ) -> Result<(), crate::Error> { let mut drain_waiters = Vec::new(); while let Some(message) = outgoing_rx.next().await { tracing::debug!(?message, "outgoing_protocol_actor"); + let (message, permit) = match message { + OutgoingMessage::Admitted { message, permit } => (*message, Some(permit)), + message => (message, None), + }; // Create the message to be sent over the transport let (json_rpc_message, destination) = match message { @@ -45,18 +92,14 @@ pub(super) async fn outgoing_protocol_actor( continue; } OutgoingMessage::BatchDispatchComplete { completion } => { - if let Some(frame) = completion.complete() { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some((frame, permit)) = completion.complete_admitted(permit) { + publish(&transport_tx, frame, permit).await?; } continue; } OutgoingMessage::BatchHandlerAttemptComplete { destination } => { - if let Some(frame) = destination.finish_handler_attempt() { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some((frame, permit)) = destination.finish_handler_attempt_admitted(permit) { + publish(&transport_tx, frame, permit).await?; } continue; } @@ -78,10 +121,8 @@ pub(super) async fn outgoing_protocol_actor( ))), ); let fallback = RawJsonRpcMessage::response(id, fallback); - if let Some(frame) = destination.abandon(fallback) { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some((frame, permit)) = destination.abandon_admitted(fallback, permit) { + publish(&transport_tx, frame, permit).await?; } continue; } @@ -99,19 +140,70 @@ pub(super) async fn outgoing_protocol_actor( continue; } - if let Some(readiness) = readiness - && let Err(error) = readiness.await - { - tracing::warn!( - ?id, - %method, - ?error, - "Outgoing request readiness failed" - ); - if let Some(pending_reply) = pending_replies.remove(&id) { - pending_reply.fail(error); + if let Some(readiness) = readiness { + enum Gate { + Ready(Result<(), crate::Error>), + Shutdown, + Urgent(OutgoingMessage), + } + let mut readiness = Box::pin(readiness); + let mut closing = Box::pin(shutdown.shutdown_requested()); + let mut skip_request = false; + loop { + let gate = future::poll_fn(|cx| { + if let Poll::Ready(result) = readiness.as_mut().poll(cx) { + return Poll::Ready(Gate::Ready(result)); + } + if closing.as_mut().poll(cx).is_ready() { + return Poll::Ready(Gate::Shutdown); + } + match outgoing_rx.poll_urgent(cx) { + Poll::Ready(Some(message)) => Poll::Ready(Gate::Urgent(message)), + _ => Poll::Pending, + } + }) + .await; + match gate { + Gate::Ready(Ok(())) => break, + Gate::Ready(Err(error)) => { + tracing::warn!(?id, %method, ?error, "Outgoing request readiness failed"); + if let Some(pending_reply) = pending_replies.remove(&id) { + pending_reply.fail(error); + } + skip_request = true; + break; + } + Gate::Shutdown => { + if let Some(pending_reply) = pending_replies.remove(&id) { + pending_reply.fail(crate::util::internal_error("connection shut down while waiting for outgoing request readiness")); + } + skip_request = true; + break; + } + Gate::Urgent(OutgoingMessage::Admitted { message, permit }) => { + if let OutgoingMessage::Notification { untyped } = *message { + publish_notification( + &transport_tx, + &protocol_compat, + &pending_replies, + untyped, + Some(permit), + ) + .await?; + } + if !pending_replies.contains(&id) { + skip_request = true; + break; + } + } + Gate::Urgent(_) => unreachable!( + "urgent admission only accepts cancellation notifications" + ), + } + } + if skip_request { + continue; } - continue; } if !pending_replies.contains(&id) { @@ -133,11 +225,13 @@ pub(super) async fn outgoing_protocol_actor( } }; - if !pending_replies.contains(&id) { + if !pending_replies.mark_published(&id) { continue; } - if let Err(error) = transport_tx.unbounded_send(TransportFrame::Single(request)) { + if let Err(error) = + publish(&transport_tx, TransportFrame::Single(request), permit).await + { let error = crate::Error::into_internal_error(error); if let Some(pending_reply) = pending_replies.remove(&id) { pending_reply.fail(error.clone()); @@ -147,32 +241,14 @@ pub(super) async fn outgoing_protocol_actor( continue; } OutgoingMessage::Notification { untyped } => { - let messages = match protocol_compat.outgoing_notification(untyped) { - Ok(messages) => messages, - Err(error) => { - tracing::warn!( - ?error, - "Dropping outgoing notification after preparation failed" - ); - continue; - } - }; - - for untyped in messages { - let message = match untyped.into_raw_jsonrpc_message(None) { - Ok(message) => message, - Err(error) => { - tracing::warn!( - ?error, - "Dropping outgoing notification after serialization failed" - ); - continue; - } - }; - transport_tx - .unbounded_send(TransportFrame::Single(message)) - .map_err(crate::Error::into_internal_error)?; - } + publish_notification( + &transport_tx, + &protocol_compat, + &pending_replies, + untyped, + permit, + ) + .await?; continue; } OutgoingMessage::Response { @@ -198,12 +274,13 @@ pub(super) async fn outgoing_protocol_actor( destination, ) } + OutgoingMessage::Admitted { .. } => { + unreachable!("application admission is unwrapped above") + } }; - if let Some(frame) = destination.complete(json_rpc_message) { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some((frame, permit)) = destination.complete_admitted(json_rpc_message, permit) { + publish(&transport_tx, frame, permit).await?; } } diff --git a/src/agent-client-protocol/src/jsonrpc/task_actor.rs b/src/agent-client-protocol/src/jsonrpc/task_actor.rs index 92a05a6d..cc30a9e7 100644 --- a/src/agent-client-protocol/src/jsonrpc/task_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/task_actor.rs @@ -1,12 +1,11 @@ use std::panic::Location; -use futures::{FutureExt, channel::mpsc, future::BoxFuture}; +use futures::{FutureExt, StreamExt, future::BoxFuture}; use crate::ConnectionTo; use crate::role::Role; -use crate::util::process_stream_concurrently; -pub type TaskTx = mpsc::UnboundedSender; +pub type TaskTx = super::admission::Sender; #[must_use] pub(crate) struct Task { @@ -54,13 +53,13 @@ impl Task { /// The "task actor" manages dynamically spawned tasks. pub(super) async fn task_actor( - task_rx: mpsc::UnboundedReceiver, + task_rx: super::admission::SimpleReceiver, _cx: &ConnectionTo, + max_running_tasks: usize, ) -> Result<(), crate::Error> { - process_stream_concurrently( - task_rx, - async |task| task.future.await, - |a, b| Box::pin(a(b)), - ) - .await + use futures::TryStreamExt as _; + task_rx + .map(Ok::<_, crate::Error>) + .try_for_each_concurrent(max_running_tasks.max(1), |task| task.future) + .await } diff --git a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs index 8737ca81..be2cb522 100644 --- a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs @@ -4,9 +4,13 @@ use std::pin::pin; use crate::jsonrpc::{RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame}; use crate::schema::v1::Response; use futures::StreamExt as _; -use futures::channel::mpsc; use serde::Deserialize as _; +/// Maximum bytes in one wire value (excluding its newline). +pub const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024; +/// Maximum number of JSON-RPC values carried in one batch. +pub const MAX_BATCH_ENTRIES: usize = 64; + enum ParsedIncomingLine { Single(RawJsonRpcMessage), Malformed { raw: String, error: crate::Error }, @@ -14,13 +18,21 @@ enum ParsedIncomingLine { } fn parse_incoming_line(line: &str) -> ParsedIncomingLine { + if line.len() > MAX_FRAME_BYTES { + return ParsedIncomingLine::Malformed { + raw: String::new(), + error: crate::Error::invalid_request().data("JSON-RPC frame exceeds maximum size"), + }; + } let value = match serde_json::from_str::(line) { Ok(value) => value, Err(error) => { tracing::debug!(?error, "Failed to parse incoming JSON-RPC JSON"); return ParsedIncomingLine::Malformed { raw: line.to_owned(), - error: crate::Error::parse_error().data(serde_json::json!({ "line": line })), + error: crate::Error::parse_error().data(serde_json::json!({ + "line": line.chars().take(256).collect::() + })), }; } }; @@ -30,6 +42,12 @@ fn parse_incoming_line(line: &str) -> ParsedIncomingLine { raw: line.to_owned(), error: crate::Error::invalid_request(), }, + serde_json::Value::Array(entries) if entries.len() > MAX_BATCH_ENTRIES => { + ParsedIncomingLine::Malformed { + raw: String::new(), + error: crate::Error::invalid_request().data("JSON-RPC batch exceeds maximum width"), + } + } serde_json::Value::Array(entries) => { let entries = entries .into_iter() @@ -60,6 +78,62 @@ fn parse_incoming_line(line: &str) -> ParsedIncomingLine { } } +/// Read newline-delimited UTF-8 without allocating an unterminated line larger +/// than the frame budget. An oversized line terminates the transport explicitly. +pub fn bounded_lines( + input: R, +) -> impl futures::Stream> { + use futures::io::BufReader; + use futures::{AsyncBufReadExt, stream}; + stream::unfold(Some(BufReader::new(input)), |reader| async move { + let mut reader = reader?; + let mut bytes = Vec::new(); + loop { + let chunk = match reader.fill_buf().await { + Ok(chunk) => chunk, + Err(error) => return Some((Err(error), None)), + }; + if chunk.is_empty() { + return if bytes.is_empty() { + None + } else { + Some(( + String::from_utf8(bytes) + .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e)), + None, + )) + }; + } + let width = chunk + .iter() + .position(|&b| b == b'\n') + .map_or(chunk.len(), |i| i + 1); + if bytes.len() + width > MAX_FRAME_BYTES + 1 { + return Some(( + Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "JSON-RPC line exceeds maximum frame size", + )), + None, + )); + } + bytes.extend_from_slice(&chunk[..width]); + reader.consume_unpin(width); + if bytes.last() == Some(&b'\n') { + bytes.pop(); + if bytes.last() == Some(&b'\r') { + bytes.pop(); + } + return Some(( + String::from_utf8(bytes) + .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e)), + Some(reader), + )); + } + } + }) +} + impl TransportFrame { /// Parse one JSON-RPC wire value while preserving batch boundaries. /// @@ -110,19 +184,21 @@ impl TransportFrame { /// /// This is the transport layer - it has no knowledge of protocol semantics (IDs, correlation, etc.). async fn transport_outgoing_frames_actor( - transport_rx: impl futures::Stream, + transport_rx: impl futures::Stream, outgoing_lines: impl futures::Sink, ) -> Result<(), crate::Error> { use futures::SinkExt; let mut transport_rx = pin!(transport_rx); let mut outgoing_lines = pin!(outgoing_lines); - while let Some(frame) = transport_rx.next().await { + while let Some(budgeted) = transport_rx.next().await { + let (frame, _permit) = budgeted.into_parts(); let json_rpc_message = match frame { TransportFrame::Single(message) => message, TransportFrame::Malformed { raw, .. } => { let raw = malformed_line_value(raw)?; tracing::trace!(message = ?raw, "Relaying invalid JSON-RPC value"); + ensure_frame_size(&raw)?; outgoing_lines .send(raw) .await @@ -133,6 +209,7 @@ async fn transport_outgoing_frames_actor( let line = serde_json::to_string(&batch).map_err(crate::Error::into_internal_error)?; tracing::trace!(message = %line, "Sending JSON-RPC batch"); + ensure_frame_size(&line)?; outgoing_lines .send(line) .await @@ -143,6 +220,7 @@ async fn transport_outgoing_frames_actor( match serde_json::to_string(&json_rpc_message) { Ok(line) => { tracing::trace!(message = %line, "Sending JSON-RPC message"); + ensure_frame_size(&line)?; outgoing_lines .send(line) .await @@ -177,6 +255,7 @@ async fn transport_outgoing_frames_actor( Err(crate::Error::internal_error()), )) .unwrap(); + ensure_frame_size(&error_line)?; outgoing_lines .send(error_line) .await @@ -189,6 +268,14 @@ async fn transport_outgoing_frames_actor( Ok(()) } +fn ensure_frame_size(line: &str) -> Result<(), crate::Error> { + if line.len() > MAX_FRAME_BYTES { + Err(crate::Error::invalid_request().data("outgoing JSON-RPC frame exceeds maximum size")) + } else { + Ok(()) + } +} + fn malformed_line_value(raw: String) -> Result { if !raw.contains('\r') && !raw.contains('\n') { return Ok(raw); @@ -202,7 +289,7 @@ fn malformed_line_value(raw: String) -> Result { } pub(super) async fn transport_outgoing_lines_actor( - transport_rx: mpsc::UnboundedReceiver, + transport_rx: super::FrameReceiver, outgoing_lines: impl futures::Sink, ) -> Result<(), crate::Error> { transport_outgoing_frames_actor(transport_rx, outgoing_lines).await @@ -222,7 +309,7 @@ pub(super) async fn transport_outgoing_lines_actor( /// This is the transport layer - it has no knowledge of protocol semantics. pub(super) async fn transport_incoming_lines_actor( incoming_lines: impl futures::Stream>, - transport_tx: mpsc::UnboundedSender, + transport_tx: super::FrameSender, ) -> Result<(), crate::Error> { let mut incoming_lines = pin!(incoming_lines); while let Some(line_result) = incoming_lines.next().await { @@ -232,17 +319,20 @@ pub(super) async fn transport_incoming_lines_actor( match parse_incoming_line(&line) { ParsedIncomingLine::Single(message) => { transport_tx - .unbounded_send(TransportFrame::Single(message)) + .send_frame(TransportFrame::Single(message)) + .await .map_err(crate::Error::into_internal_error)?; } ParsedIncomingLine::Malformed { raw, error } => { transport_tx - .unbounded_send(TransportFrame::Malformed { raw, error }) + .send_frame(TransportFrame::Malformed { raw, error }) + .await .map_err(crate::Error::into_internal_error)?; } ParsedIncomingLine::Batch(entries) => { transport_tx - .unbounded_send(TransportFrame::Batch(entries)) + .send_frame(TransportFrame::Batch(entries)) + .await .map_err(crate::Error::into_internal_error)?; } } @@ -257,6 +347,28 @@ mod tests { use super::*; use crate::ErrorCode; + #[test] + fn rejects_batches_over_width_limit() { + let batch = format!("[{}]", vec!["null"; MAX_BATCH_ENTRIES + 1].join(",")); + let ParsedIncomingLine::Malformed { error, .. } = parse_incoming_line(&batch) else { + panic!("oversized batch must be rejected"); + }; + assert_eq!(error.code, ErrorCode::InvalidRequest); + } + + #[tokio::test] + async fn oversized_unterminated_line_fails_before_eof() { + let input = futures::io::Cursor::new(vec![b'x'; MAX_FRAME_BYTES + 2]); + let mut lines = Box::pin(bounded_lines(input)); + let error = lines + .next() + .await + .expect("explicit framing failure") + .unwrap_err(); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + assert!(lines.next().await.is_none()); + } + #[test] fn parses_batch_entries_independently() { let ParsedIncomingLine::Batch(batch) = parse_incoming_line( @@ -446,15 +558,19 @@ mod tests { Ok::<_, std::io::Error>(captured) }); - transport_outgoing_frames_actor( - futures::stream::iter([TransportFrame::Malformed { + let (source, destination) = crate::Channel::duplex(); + source + .tx + .send_frame(TransportFrame::Malformed { raw: raw.clone(), error: crate::Error::parse_error(), - }]), - outgoing, - ) - .await - .unwrap(); + }) + .await + .unwrap(); + drop(source); + transport_outgoing_frames_actor(destination.rx, outgoing) + .await + .unwrap(); let lines = captured.lock().unwrap(); assert_eq!(lines.len(), 1); diff --git a/src/agent-client-protocol/src/lib.rs b/src/agent-client-protocol/src/lib.rs index 0a543aed..6d280fc2 100644 --- a/src/agent-client-protocol/src/lib.rs +++ b/src/agent-client-protocol/src/lib.rs @@ -143,12 +143,13 @@ pub mod util; pub use capabilities::*; pub use jsonrpc::{ - Builder, ByteStreams, Channel, ConnectionContext, ConnectionTo, Dispatch, DynamicHandlerGuard, - HandleConnectionClose, HandleDispatchFrom, Handled, INCOMING_TRANSPORT_CLOSED_REASON, - IntoHandled, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, Lines, - NullClose, NullHandler, RawConnectionContext, RawJsonRpcMessage, RawJsonRpcParams, Responder, - ResponseRouter, SentRequest, TransportBatch, TransportBatchEntry, TransportFrame, - UntypedMessage, is_incoming_transport_closed, + BudgetedFrame, Builder, ByteStreams, Channel, ConnectionContext, ConnectionLimits, + ConnectionTo, Dispatch, DynamicHandlerGuard, FrameAdmission, FramePermit, FrameReceiver, + FrameSender, HandleConnectionClose, HandleDispatchFrom, Handled, + INCOMING_TRANSPORT_CLOSED_REASON, IntoHandled, JsonRpcMessage, JsonRpcNotification, + JsonRpcRequest, JsonRpcResponse, Lines, NullClose, NullHandler, RawConnectionContext, + RawJsonRpcMessage, RawJsonRpcParams, Responder, ResponseRouter, SentRequest, TransportBatch, + TransportBatchEntry, TransportFrame, UntypedMessage, is_incoming_transport_closed, run::{ChainRun, NullRun, RunWithConnectionTo}, }; pub use jsonrpc::{RequestCancellation, is_cancel_request_notification}; @@ -162,7 +163,7 @@ pub use role::{ acp::{Agent, Client, Conductor, Proxy}, }; -pub use component::{ConnectTo, DynConnectTo}; +pub use component::{ConnectTo, ConnectionDriver, DynConnectTo}; /// Implementation details used by the derive macros. #[doc(hidden)] diff --git a/src/agent-client-protocol/src/mcp_server/active_session.rs b/src/agent-client-protocol/src/mcp_server/active_session.rs index a5e1ce0f..0d3ed260 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -1,4 +1,4 @@ -//! Request-scoped native MCP transport. An ACP request owns exactly one backend instance. +//! Request-scoped native MCP transport. Each ACP request owns execution and cleanup. use futures::{ StreamExt, @@ -15,19 +15,21 @@ use std::{ use crate::{ Agent, Channel, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, RawJsonRpcMessage, Responder, Role, - TransportFrame, - mcp_server::{McpConnectionContext, McpConnectionTo, McpServerConnect}, + JsonRpcNotification, JsonRpcRequest, RawJsonRpcMessage, Responder, Role, TransportFrame, + mcp_server::{ + MCP_BACKEND_FAILURE, MCP_RESOURCE_EXHAUSTED, McpConnectionContext, McpConnectionTo, + McpOperationCancellation, McpOutcome, McpRequest, McpRequestContext, McpServerConnect, + McpService, + }, role::HasPeer, schema::v1::{ - McpRequestId, McpServerAcpId, MessageMcpNotification, MessageMcpRequest, + McpError, McpRequestId, McpServerAcpId, MessageMcpNotification, MessageMcpRequest, MessageMcpResponse, RequestId, }, util::MatchDispatchFrom, }; -// These bound admitted work and individual payloads, not the SDK's underlying -// Channel/outgoing queues. End-to-end native backpressure is separate transport work. +// These bound admitted work and individual payloads. const MAX_ACTIVE_REQUESTS: usize = 64; const MAX_PAYLOAD_BYTES: usize = 16 * 1024 * 1024; const MCP_VERSION: &str = "2026-07-28"; @@ -38,8 +40,7 @@ pub(super) struct V1McpProtocol; pub(super) struct V2McpProtocol; pub(super) trait McpProtocol: Send + 'static { - type MessageRequest: JsonRpcRequest; - type MessageResponse: JsonRpcResponse; + type MessageRequest: JsonRpcRequest; type MessageNotification: JsonRpcNotification; fn server_id(request: &Self::MessageRequest) -> McpServerAcpId; @@ -55,7 +56,6 @@ pub(super) trait McpProtocol: Send + 'static { impl McpProtocol for V1McpProtocol { type MessageRequest = MessageMcpRequest; - type MessageResponse = MessageMcpResponse; type MessageNotification = MessageMcpNotification; fn server_id(request: &Self::MessageRequest) -> McpServerAcpId { @@ -77,10 +77,56 @@ impl McpProtocol for V1McpProtocol { } } +fn into_mcp_error(error: crate::Error) -> McpError { + let mut mcp = McpError::new(error.code.into(), error.message); + if let Some(data) = error.data { + mcp = mcp.data(data); + } + mcp +} + +fn outcome_response(outcome: McpOutcome) -> Result { + let response = match outcome { + McpOutcome::Result(value) => MessageMcpResponse::success(value), + McpOutcome::Error(error) => MessageMcpResponse::error(error), + }; + check_payload_size(&response, MAX_PAYLOAD_BYTES)?; + Ok(response) +} + +fn project_outcome(outcome: McpOutcome, is_discovery: bool) -> Result { + match outcome { + McpOutcome::Result(mut value) if is_discovery => { + constrain_discovery_versions(&mut value)?; + Ok(McpOutcome::Result(value)) + } + other => Ok(other), + } +} + +fn send_outcome( + responder: Responder, + result: Result, + is_discovery: bool, +) -> Result<(), crate::Error> { + match result { + Ok(outcome) => { + // Projection failures are MCP outcomes; binding and size failures + // remain named outer ACP errors, regardless of backend type. + let outcome = project_outcome(outcome, is_discovery) + .unwrap_or_else(|error| McpOutcome::Error(into_mcp_error(error))); + match outcome_response(outcome) { + Ok(response) => responder.respond(response), + Err(error) => responder.respond_with_error(error), + } + } + Err(error) => responder.respond_with_error(error), + } +} + #[cfg(feature = "unstable_protocol_v2")] impl McpProtocol for V2McpProtocol { type MessageRequest = crate::schema::v2::MessageMcpRequest; - type MessageResponse = crate::schema::v2::MessageMcpResponse; type MessageNotification = crate::schema::v2::MessageMcpNotification; fn server_id(request: &Self::MessageRequest) -> McpServerAcpId { @@ -107,6 +153,7 @@ impl McpProtocol for V2McpProtocol { pub(super) struct McpActiveSession { server_id: McpServerAcpId, mcp_connect: Arc>, + service: Option>>, active: ActiveRequests, protocol: PhantomData Protocol>, } @@ -138,7 +185,7 @@ fn admit_request( } if requests.len() >= MAX_ACTIVE_REQUESTS { return Err( - crate::Error::new(-32000, "MCP active request limit exceeded") + crate::Error::new(MCP_RESOURCE_EXHAUSTED, "MCP active request limit exceeded") .data(serde_json::json!({"limit": MAX_ACTIVE_REQUESTS})), ); } @@ -168,7 +215,7 @@ fn check_payload_size(value: &impl serde::Serialize, limit: usize) -> Result<(), } } serde_json::to_writer(Budget(limit), value).map_err(|_| { - crate::Error::new(-32000, "MCP payload limit exceeded") + crate::Error::new(MCP_RESOURCE_EXHAUSTED, "MCP payload limit exceeded") .data(serde_json::json!({"limitBytes": limit})) }) } @@ -178,13 +225,15 @@ where Counterpart: HasPeer, Protocol: McpProtocol, { - pub fn new( + pub fn new_with_service( server_id: McpServerAcpId, mcp_connect: Arc>, + service: Option>>, ) -> Self { Self { server_id, mcp_connect, + service, active: Arc::default(), protocol: PhantomData, } @@ -193,15 +242,10 @@ where fn handle_request( &mut self, request: Protocol::MessageRequest, - responder: Responder, + responder: Responder, connection: &ConnectionTo, - ) -> Result< - Handled<( - Protocol::MessageRequest, - Responder, - )>, - crate::Error, - > { + ) -> Result)>, crate::Error> + { let server_id = Protocol::server_id(&request); if server_id != self.server_id { return Ok(Handled::No { @@ -211,12 +255,15 @@ where } let request_id = Protocol::request_id(&request); let (method, params) = Protocol::into_request(request); - if let Err(error) = validate_modern_request(&method, params.as_ref()) - .and_then(|()| check_payload_size(&(&method, ¶ms, &request_id), MAX_PAYLOAD_BYTES)) + if let Err(error) = check_payload_size(&(&method, ¶ms, &request_id), MAX_PAYLOAD_BYTES) { responder.respond_with_error(error)?; return Ok(Handled::Yes); } + if let Err(error) = validate_modern_request(&method, params.as_ref()) { + responder.respond(outcome_response(McpOutcome::Error(into_mcp_error(error)))?)?; + return Ok(Handled::Yes); + } let (guard, stop_rx) = match admit_request(&self.active, request_id.clone()) { Ok(admitted) => admitted, Err(error) => { @@ -225,30 +272,143 @@ where } }; - let backend = self.mcp_connect.connect(McpConnectionTo { + if let Some(service) = self.service.clone() { + let metadata = params + .as_ref() + .and_then(|params| params.get("_meta")) + .and_then(Value::as_object) + .expect("validated MCP metadata") + .clone(); + let cancellation = responder.cancellation(); + let operation_cancellation = McpOperationCancellation::new(); + let alive = Arc::new(futures::lock::Mutex::new(true)); + let send_connection = connection.clone(); + let send_server_id = server_id.clone(); + let send_request_id = request_id.clone(); + let send_cancellation = cancellation.clone(); + let send_operation_cancellation = operation_cancellation.clone(); + let send_alive = alive.clone(); + let notify = Arc::new(move |method: String, params: Option>| { + let connection = send_connection.clone(); + let server_id = send_server_id.clone(); + let request_id = send_request_id.clone(); + let cancellation = send_cancellation.clone(); + let operation_cancellation = send_operation_cancellation.clone(); + let alive = send_alive.clone(); + let send = async move { + let active = alive.lock().await; + if !*active + || cancellation.is_cancelled() + || operation_cancellation.is_cancelled() + { + return Err(crate::Error::request_cancelled()); + } + check_payload_size(&(&method, ¶ms), MAX_PAYLOAD_BYTES)?; + let send = connection.send_notification_to_async( + Agent, + Protocol::notification(server_id, request_id, method, params), + ); + futures::pin_mut!(send); + let cancelled = async { + let peer = cancellation.cancelled(); + let operation = operation_cancellation.cancelled(); + let shutdown = connection.shutdown_requested(); + futures::pin_mut!(peer, operation, shutdown); + let peer_or_operation = future::select(peer, operation); + futures::pin_mut!(peer_or_operation); + let _reason = future::select(peer_or_operation, shutdown).await; + }; + futures::pin_mut!(cancelled); + let result = match future::select(send, cancelled).await { + Either::Left((result, _)) => result, + Either::Right(((), _)) => Err(crate::Error::request_cancelled()), + }; + drop(active); + result + }; + Box::pin(send) as futures::future::BoxFuture<'static, Result<(), crate::Error>> + }); + let context = McpRequestContext::new( + server_id.clone(), + request_id.clone(), + McpConnectionTo { + context: McpConnectionContext::Acp { + server_id, + request_id, + }, + connection: connection.clone(), + cleanup: Some(Arc::default()), + }, + metadata, + cancellation.clone(), + operation_cancellation.clone(), + notify, + ); + let is_discovery = method == "server/discover"; + let shutdown_connection = connection.clone(); + connection.spawn(async move { + let request = McpRequest { method, params }; + let cleanup_connection = context.connection().clone(); + let operation = service.execute(request, context); + let stop = async { + let cancelled = cancellation.cancelled(); + let shutdown = shutdown_connection.shutdown_requested(); + futures::pin_mut!(cancelled); + futures::pin_mut!(shutdown); + let stop_rx = stop_rx; + futures::pin_mut!(stop_rx); + let cancel_or_shutdown = future::select(cancelled, shutdown); + futures::pin_mut!(cancel_or_shutdown); + let _reason = future::select(cancel_or_shutdown, stop_rx).await; + }; + let result = match future::select(operation, Box::pin(stop)).await { + Either::Left((result, _)) => result, + Either::Right(((), operation)) => { + operation_cancellation.cancel(); + *alive.lock().await = false; + // Do not discard the operation future: its completion + // includes rmcp handler cancellation and actor join. + drop(operation.await); + Err(crate::Error::request_cancelled()) + } + }; + *alive.lock().await = false; + cleanup_connection.wait_cleanup().await; + // Operation futures have been dropped and cannot send late output. + drop(guard); + let response = send_outcome(responder, result, is_discovery); + if let Err(error) = response { + tracing::debug!(?error, "cannot send MCP response"); + } + Ok(()) + })?; + return Ok(Handled::Yes); + } + + let cleanup_connection = McpConnectionTo { context: McpConnectionContext::Acp { server_id: server_id.clone(), request_id: request_id.clone(), }, connection: connection.clone(), - }); + cleanup: Some(Arc::default()), + }; + let backend = self.mcp_connect.connect(cleanup_connection.clone()); let connection_for_task = connection.clone(); let cancellation = responder.cancellation(); let (mut client, server) = Channel::duplex(); - // Dropping this sender when the request completes stops the backend even if it - // has outstanding work after emitting its final response. + // Keep the operation admitted until its backend has actually stopped. let (backend_stop_tx, backend_stop_rx) = oneshot::channel::<()>(); + let (backend_done_tx, mut backend_done_rx) = oneshot::channel(); let spawn_result = connection.spawn(async move { - let run = backend.connect_to(server); - futures::pin_mut!(run); - let stop = backend_stop_rx; - futures::pin_mut!(stop); - match future::select(run, stop).await { - Either::Left((Err(error), _)) => { - tracing::warn!(?error, "request-scoped MCP backend failed"); - } - Either::Left((Ok(()), _)) | Either::Right((_, _)) => {} - } + // Own (not merely borrow) the future so cancellation drops its + // backend before the completion acknowledgement is published. + let run = Box::pin(backend.connect_to(server)); + let outcome = match future::select(run, backend_stop_rx).await { + Either::Left((result, _)) => result, + Either::Right((_, _)) => Ok(()), + }; + drop(backend_done_tx.send(outcome)); Ok(()) }); if let Err(error) = spawn_result { @@ -267,62 +427,73 @@ where )?; client .tx - .unbounded_send(TransportFrame::Single(raw)) + .send_frame(TransportFrame::Single(raw)) + .await .map_err(crate::Error::into_internal_error)?; - while let Some(frame) = client.rx.next().await { + while let Some(budgeted) = client.rx.next().await { + let (frame, _permit) = budgeted.into_parts(); let TransportFrame::Single(message) = frame else { - return Err(crate::Error::invalid_request() - .data("MCP backends must send individual valid JSON-RPC messages")); + return Err(crate::Error::new( + MCP_BACKEND_FAILURE, + "MCP backends must send individual valid JSON-RPC messages", + )); }; if matches!(message, RawJsonRpcMessage::Response(_)) && message.response_id() != Some(&inner_id) { - return Err(crate::Error::invalid_params() - .data("MCP backend returned a different request ID")); + return Err(crate::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend returned a different request ID", + )); } match message { RawJsonRpcMessage::Response(response) => { - check_payload_size(&response, MAX_PAYLOAD_BYTES)?; // Returning ends notification forwarding before the terminal reply. return match response { - crate::schema::v1::Response::Result { mut result, .. } => { - if is_discovery { - constrain_discovery_versions(&mut result)?; - } - Ok(result) + crate::schema::v1::Response::Result { result, .. } => { + Ok(McpOutcome::Result(result)) + } + crate::schema::v1::Response::Error { error, .. } => { + Ok(McpOutcome::Error(into_mcp_error(error))) } - crate::schema::v1::Response::Error { error, .. } => Err(error), }; } RawJsonRpcMessage::Notification(notification) => { check_payload_size(¬ification, MAX_PAYLOAD_BYTES)?; - let params = - match notification.params { - Some(params) => match params.into_value() { - Value::Object(map) => Some(map), - _ => return Err(crate::Error::invalid_params().data( + let params = match notification.params { + Some(params) => match params.into_value() { + Value::Object(map) => Some(map), + _ => { + return Err(crate::Error::new( + MCP_BACKEND_FAILURE, "MCP backend notification parameters must be an object", - )), - }, - None => None, - }; - connection_for_task.send_notification_to( - Agent, - Protocol::notification( - server_id.clone(), - request_id.clone(), - notification.method.to_string(), - params, - ), - )?; + )); + } + }, + None => None, + }; + connection_for_task + .send_notification_to_async( + Agent, + Protocol::notification( + server_id.clone(), + request_id.clone(), + notification.method.to_string(), + params, + ), + ) + .await?; } RawJsonRpcMessage::Request(_) => { - return Err(crate::Error::method_not_found() - .data("reverse MCP requests are not supported")); + return Err(crate::Error::new( + MCP_BACKEND_FAILURE, + "reverse MCP requests are not supported", + )); } } } - Err(crate::util::internal_error( + Err(crate::Error::new( + MCP_BACKEND_FAILURE, "MCP backend closed without a response", )) }; @@ -330,26 +501,39 @@ where .run_until_cancelled(async { let process = process; futures::pin_mut!(process); - let stop = stop_rx; + let stop = async { + let _reason = future::select( + stop_rx, + Box::pin(connection_for_task.shutdown_requested()), + ) + .await; + }; futures::pin_mut!(stop); - match future::select(process, stop).await { + let work = async { + match future::select(process, stop).await { + Either::Left((result, _)) => result, + Either::Right(((), _)) => Err(crate::Error::request_cancelled()), + } + }; + futures::pin_mut!(work); + match future::select(work, &mut backend_done_rx).await { Either::Left((result, _)) => result, - Either::Right((_, _)) => Err(crate::Error::request_cancelled()), + Either::Right((Ok(Err(error)), _)) => Err(error), + // The backend can finish immediately after queueing its + // reply. Drain the channel before calling that an EOF. + Either::Right((Ok(Ok(())) | Err(_), work)) => work.await, } }) .await; - // No more notifications can be forwarded after `process` is dropped. - // Release the ID before publishing the final response so a caller can - // immediately reuse it for the next independent operation. + // Revoking the channel stops any late output. A cancellation is only + // caller-visible now; cleanup and ID release happen after backend exit. drop(backend_stop_tx); + // The receiver can have already completed in the race above. Polling + // it again then returns immediately; otherwise this joins cleanup. + drop(backend_done_rx.await); + cleanup_connection.wait_cleanup().await; drop(guard); - let response = match result { - Ok(value) => match Protocol::MessageResponse::from_value("mcp/message", value) { - Ok(response) => responder.respond(response), - Err(error) => responder.respond_with_error(error), - }, - Err(error) => responder.respond_with_error(error), - }; + let response = send_outcome(responder, result, is_discovery); if let Err(error) = response { tracing::debug!(?error, "cannot send request-scoped MCP response"); } @@ -445,12 +629,34 @@ fn validate_modern_request( #[cfg(test)] mod tests { use super::{ - ActiveRequests, MAX_ACTIVE_REQUESTS, admit_request, check_payload_size, - constrain_discovery_versions, validate_modern_request, + ActiveRequests, MAX_ACTIVE_REQUESTS, MAX_PAYLOAD_BYTES, McpOutcome, admit_request, + check_payload_size, constrain_discovery_versions, into_mcp_error, outcome_response, + validate_modern_request, + }; + use crate::{ + mcp_server::MCP_RESOURCE_EXHAUSTED, + schema::v1::{McpError, McpRequestId}, }; - use crate::schema::v1::McpRequestId; use serde_json::json; + #[test] + fn both_outcome_branches_obey_the_binding_payload_limit() { + for outcome in [ + McpOutcome::Result(json!("x".repeat(MAX_PAYLOAD_BYTES))), + McpOutcome::Error( + McpError::new(-32000, "peer error").data(json!("x".repeat(MAX_PAYLOAD_BYTES))), + ), + ] { + let error = outcome_response(outcome).expect_err("oversized carrier must be rejected"); + assert_eq!(i32::from(error.code), MCP_RESOURCE_EXHAUSTED); + } + let result = outcome_response(McpOutcome::Result(serde_json::Value::Null)).unwrap(); + assert_eq!( + serde_json::to_value(result).unwrap(), + json!({"result":null}) + ); + } + #[test] fn only_modern_request_metadata_is_accepted() { let modern = json!({"_meta": {"io.modelcontextprotocol/protocolVersion": "2026-07-28", "io.modelcontextprotocol/clientCapabilities": {}, "requestState": {"opaque": true}}}); @@ -502,7 +708,7 @@ mod tests { let overload = admit_request(&active, McpRequestId::new("extra")) .err() .unwrap(); - assert_eq!(i32::from(overload.code), -32000); + assert_eq!(i32::from(overload.code), MCP_RESOURCE_EXHAUSTED); drop(admitted.pop()); let replacement = admit_request(&active, McpRequestId::new("replacement")).unwrap(); assert_eq!(active.lock().unwrap().len(), MAX_ACTIVE_REQUESTS); @@ -535,4 +741,30 @@ mod tests { assert!(constrain_discovery_versions(&mut unsupported).is_err()); assert!(constrain_discovery_versions(&mut json!({})).is_err()); } + + #[test] + fn mcp_validation_errors_use_inner_carrier_and_preserve_null_data() { + let unsupported = validate_modern_request( + "tools/list", + json!({"_meta": { + "io.modelcontextprotocol/protocolVersion": "2025-03-26", + "io.modelcontextprotocol/clientCapabilities": {} + }}) + .as_object(), + ) + .expect_err("unsupported inner version"); + let response = outcome_response(McpOutcome::Error(into_mcp_error(unsupported))).unwrap(); + assert_eq!( + serde_json::to_value(response).unwrap()["error"]["code"], + -32022 + ); + let response = outcome_response(McpOutcome::Error( + McpError::new(-32000, "opaque MCP error").data(serde_json::Value::Null), + )) + .unwrap(); + assert_eq!( + serde_json::to_value(response).unwrap(), + json!({"error": {"code": -32000, "message": "opaque MCP error", "data": null}}) + ); + } } diff --git a/src/agent-client-protocol/src/mcp_server/context.rs b/src/agent-client-protocol/src/mcp_server/context.rs index 517f3026..876390d4 100644 --- a/src/agent-client-protocol/src/mcp_server/context.rs +++ b/src/agent-client-protocol/src/mcp_server/context.rs @@ -1,4 +1,8 @@ use crate::{ConnectionTo, role::Role}; +#[cfg(feature = "unstable_mcp_over_acp")] +use futures::channel::oneshot; +#[cfg(feature = "unstable_mcp_over_acp")] +use std::sync::{Arc, Mutex}; #[cfg(feature = "unstable_mcp_over_acp")] use crate::schema::v1::{McpRequestId, McpServerAcpId}; @@ -58,9 +62,28 @@ impl McpConnectionContext { pub struct McpConnectionTo { pub(super) context: McpConnectionContext, pub(super) connection: ConnectionTo, + #[cfg(feature = "unstable_mcp_over_acp")] + pub(super) cleanup: Option>>>>, } impl McpConnectionTo { + #[cfg(feature = "unstable_mcp_over_acp")] + pub(crate) fn register_cleanup(&self, done: oneshot::Receiver<()>) { + if let Some(cleanup) = &self.cleanup { + cleanup.lock().expect("MCP cleanup poisoned").push(done); + } + } + + #[cfg(feature = "unstable_mcp_over_acp")] + pub(crate) async fn wait_cleanup(&self) { + if let Some(cleanup) = &self.cleanup { + let pending = std::mem::take(&mut *cleanup.lock().expect("MCP cleanup poisoned")); + for done in pending { + let _ = done.await; + } + } + } + /// Describes whether this is a standalone or ACP-attached MCP connection. #[must_use] pub fn context(&self) -> &McpConnectionContext { diff --git a/src/agent-client-protocol/src/mcp_server/mod.rs b/src/agent-client-protocol/src/mcp_server/mod.rs index a77b53a1..e19d2d94 100644 --- a/src/agent-client-protocol/src/mcp_server/mod.rs +++ b/src/agent-client-protocol/src/mcp_server/mod.rs @@ -57,6 +57,8 @@ mod context; #[cfg(feature = "schemars")] mod registry; mod server; +#[cfg(feature = "unstable_mcp_over_acp")] +mod service; #[cfg(feature = "schemars")] mod tool; #[cfg(feature = "schemars")] @@ -70,9 +72,20 @@ pub use registry::{ EnabledTools, McpToolMetadata, McpToolRegistry, McpToolSchema, RegisteredMcpTool, }; pub use server::McpServer; +#[cfg(feature = "unstable_mcp_over_acp")] +pub use service::{ + McpOperationCancellation, McpOutcome, McpRequest, McpRequestContext, McpService, +}; #[cfg(feature = "schemars")] #[cfg_attr(docsrs, doc(cfg(feature = "schemars")))] pub use tool::McpTool; #[cfg(feature = "schemars")] #[cfg_attr(docsrs, doc(cfg(feature = "schemars")))] pub use tool_fn::{tool_fn, tool_fn_mut}; + +/// ACP binding error: the MCP operation admission or payload budget was exhausted. +pub const MCP_RESOURCE_EXHAUSTED: i32 = -33000; +/// ACP binding error: the requested MCP server registration is no longer available. +pub const MCP_SERVER_UNAVAILABLE: i32 = -33001; +/// ACP binding error: the MCP backend failed before returning an MCP outcome. +pub const MCP_BACKEND_FAILURE: i32 = -33002; diff --git a/src/agent-client-protocol/src/mcp_server/server.rs b/src/agent-client-protocol/src/mcp_server/server.rs index 2b34859f..5f86ee18 100644 --- a/src/agent-client-protocol/src/mcp_server/server.rs +++ b/src/agent-client-protocol/src/mcp_server/server.rs @@ -4,6 +4,8 @@ use std::{marker::PhantomData, sync::Arc}; use futures::{StreamExt, channel::mpsc}; +#[cfg(feature = "unstable_mcp_over_acp")] +use crate::mcp_server::McpService; use crate::{ ConnectTo, Dispatch, DynConnectTo, Role, jsonrpc::run::{NullRun, RunWithConnectionTo}, @@ -66,6 +68,8 @@ pub struct McpServer { /// The "connect" instance connect: Arc>, + #[cfg(feature = "unstable_mcp_over_acp")] + service: Option>>, /// The runner is a task that should be run alongside the message handler. /// Some futures direct messages back through channels to this future which actually @@ -103,6 +107,37 @@ where McpServer { phantom: PhantomData, connect: Arc::new(c), + #[cfg(feature = "unstable_mcp_over_acp")] + service: None, + runner, + } + } + + /// Construct a reusable request-native application service for ACP. + /// + /// Standalone serving additionally needs a direct MCP transport adapter; + /// use [`Self::new_service_with_standalone`] when direct serving is needed. + #[cfg(feature = "unstable_mcp_over_acp")] + pub fn new_service( + service: impl McpService, + name: impl Into, + runner: Run, + ) -> Self { + Self::new_service_with_standalone(service, NoStandalone { name: name.into() }, runner) + } + + /// Construct a reusable application service with a separate standalone MCP + /// connector. ACP requests never create connector sessions. + #[cfg(feature = "unstable_mcp_over_acp")] + pub fn new_service_with_standalone( + service: impl McpService, + standalone: impl McpServerConnect, + runner: Run, + ) -> Self { + Self { + phantom: PhantomData, + connect: Arc::new(standalone), + service: Some(Arc::new(service)), runner, } } @@ -116,10 +151,14 @@ where let Self { phantom: _, connect, + service, runner, } = self; let server_id = McpServerAcpId::new(format!("mcp-server:{}", Uuid::new_v4())); - (McpSessionHandler::new(server_id, connect), runner) + ( + McpSessionHandler::new_with_service(server_id, connect, service), + runner, + ) } /// Split this MCP server into a protocol v2 session handler and its runner. @@ -131,10 +170,40 @@ where let Self { phantom: _, connect, + service, runner, } = self; let server_id = McpServerAcpId::new(format!("mcp-server:{}", Uuid::new_v4())); - (V2McpSessionHandler::new(server_id, connect), runner) + ( + V2McpSessionHandler::new_with_service(server_id, connect, service), + runner, + ) + } +} + +#[cfg(feature = "unstable_mcp_over_acp")] +struct NoStandalone { + name: String, +} + +#[cfg(feature = "unstable_mcp_over_acp")] +impl McpServerConnect for NoStandalone { + fn name(&self) -> String { + self.name.clone() + } + + fn connect(&self, _context: McpConnectionTo) -> DynConnectTo { + struct Unavailable; + impl ConnectTo for Unavailable { + fn connect_to( + self, + _client: impl ConnectTo, + ) -> impl Future> + Send { + std::future::ready(Err(crate::Error::method_not_found() + .data("this MCP service has no standalone transport adapter"))) + } + } + DynConnectTo::new(Unavailable) } } @@ -154,9 +223,17 @@ impl McpSessionHandler where Counterpart: HasPeer, { - pub fn new(server_id: McpServerAcpId, connect: Arc>) -> Self { + fn new_with_service( + server_id: McpServerAcpId, + connect: Arc>, + service: Option>>, + ) -> Self { Self { - active_session: McpActiveSession::new(server_id.clone(), connect.clone()), + active_session: McpActiveSession::new_with_service( + server_id.clone(), + connect.clone(), + service, + ), server_id, connect, } @@ -187,9 +264,22 @@ impl V2McpSessionHandler where Counterpart: HasPeer, { + #[cfg(test)] fn new(server_id: McpServerAcpId, connect: Arc>) -> Self { + Self::new_with_service(server_id, connect, None) + } + + fn new_with_service( + server_id: McpServerAcpId, + connect: Arc>, + service: Option>>, + ) -> Self { Self { - active_session: McpActiveSession::new(server_id.clone(), connect.clone()), + active_session: McpActiveSession::new_with_service( + server_id.clone(), + connect.clone(), + service, + ), server_id, connect, } @@ -405,6 +495,8 @@ where connect, runner, phantom: _, + #[cfg(feature = "unstable_mcp_over_acp")] + service: _, } = self; let (tx, mut rx) = mpsc::unbounded(); @@ -424,6 +516,8 @@ where connect.connect(McpConnectionTo { context: McpConnectionContext::Standalone, connection: connection_to_client.clone(), + #[cfg(feature = "unstable_mcp_over_acp")] + cleanup: None, }); role::mcp::Client diff --git a/src/agent-client-protocol/src/mcp_server/service.rs b/src/agent-client-protocol/src/mcp_server/service.rs new file mode 100644 index 00000000..30eee4a5 --- /dev/null +++ b/src/agent-client-protocol/src/mcp_server/service.rs @@ -0,0 +1,197 @@ +//! Request-native application services for MCP-over-ACP. + +use std::sync::Arc; + +use futures::{ + channel::oneshot, + future::{BoxFuture, FutureExt, Shared}, +}; +use serde_json::{Map, Value}; + +use super::McpConnectionTo; +use crate::{ + Error, RequestCancellation, Role, + schema::v1::{McpError, McpRequestId, McpServerAcpId}, +}; + +/// One MCP invocation. Its application service may be reused across invocations. +#[derive(Debug)] +pub struct McpRequest { + /// The MCP method. + pub method: String, + /// Its MCP parameters; metadata is validated before dispatch. + pub params: Option>, +} + +/// The MCP outcome is distinct from a failure in the ACP binding itself. +#[derive(Debug)] +pub enum McpOutcome { + /// Successful, opaque MCP result. + Result(Value), + /// An unmodified MCP error object (including optional or explicitly null data). + Error(McpError), +} + +type Notify = dyn Fn(String, Option>) -> BoxFuture<'static, Result<(), Error>> + + Send + + Sync; + +/// Explicit cancellation of an operation, including provider removal and +/// connection shutdown (which need not cancel the original ACP request). +#[derive(Clone)] +pub struct McpOperationCancellation { + state: Arc, +} + +impl std::fmt::Debug for McpOperationCancellation { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("McpOperationCancellation") + .field("cancelled", &self.is_cancelled()) + .finish() + } +} + +struct CancellationState { + cancelled: std::sync::atomic::AtomicBool, + sender: std::sync::Mutex>>, + signal: Shared>, +} + +impl McpOperationCancellation { + pub(crate) fn new() -> Self { + let (tx, rx) = oneshot::channel(); + Self { + state: Arc::new(CancellationState { + cancelled: std::sync::atomic::AtomicBool::new(false), + sender: std::sync::Mutex::new(Some(tx)), + signal: rx.map(|_| ()).boxed().shared(), + }), + } + } + + pub(crate) fn cancel(&self) { + self.state + .cancelled + .store(true, std::sync::atomic::Ordering::Release); + drop( + self.state + .sender + .lock() + .expect("MCP cancellation poisoned") + .take(), + ); + } + + /// Await cancellation from the caller, provider, or transport. + pub async fn cancelled(&self) { + self.state.signal.clone().await; + } + /// Whether the operation may still produce output. + #[must_use] + pub fn is_cancelled(&self) -> bool { + self.state + .cancelled + .load(std::sync::atomic::Ordering::Acquire) + } +} + +/// Per-operation authority. Notifications are admitted only while this request +/// is live; retaining the service does not retain an operation's output rights. +#[derive(Clone)] +pub struct McpRequestContext { + server_id: McpServerAcpId, + request_id: McpRequestId, + connection: McpConnectionTo, + metadata: Map, + cancellation: RequestCancellation, + operation_cancellation: McpOperationCancellation, + notify: Arc, +} + +impl std::fmt::Debug for McpRequestContext { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("McpRequestContext") + .field("server_id", &self.server_id) + .field("request_id", &self.request_id) + .field("metadata", &self.metadata) + .field("operation_cancellation", &self.operation_cancellation) + .finish_non_exhaustive() + } +} + +impl McpRequestContext { + pub(crate) fn new( + server_id: McpServerAcpId, + request_id: McpRequestId, + connection: McpConnectionTo, + metadata: Map, + cancellation: RequestCancellation, + operation_cancellation: McpOperationCancellation, + notify: Arc, + ) -> Self { + Self { + server_id, + request_id, + connection, + metadata, + cancellation, + operation_cancellation, + notify, + } + } + + /// Server identifier bound to this operation. + pub fn server_id(&self) -> &McpServerAcpId { + &self.server_id + } + /// Logical operation identifier. + pub fn request_id(&self) -> &McpRequestId { + &self.request_id + } + /// Host connection, available to application tools. + pub fn connection(&self) -> &McpConnectionTo { + &self.connection + } + /// Validated MCP metadata, including the negotiated protocol version and + /// the client's capability declaration. + pub fn metadata(&self) -> &Map { + &self.metadata + } + /// Request cancellation handle. + pub fn cancellation(&self) -> &RequestCancellation { + &self.cancellation + } + /// Cancellation for this operation, including provider removal and EOF. + pub fn operation_cancellation(&self) -> &McpOperationCancellation { + &self.operation_cancellation + } + + /// Send a bounded, operation-scoped MCP notification. + pub async fn send_notification( + &self, + method: impl Into, + params: Option>, + ) -> Result<(), Error> { + if self.cancellation.is_cancelled() || self.operation_cancellation.is_cancelled() { + return Err(Error::request_cancelled()); + } + (self.notify)(method.into(), params).await + } +} + +/// Reusable application service. An invocation owns its returned future; an +/// implementation may deliberately share application state between requests. +pub trait McpService: Send + Sync + 'static { + /// Execute one MCP request, returning an owned operation future. + /// + /// The future includes backend teardown: on + /// [`McpRequestContext::operation_cancellation`], stop user work and finish + /// owned cleanup before returning. The binding keeps admission until this + /// future completes rather than abandoning cleanup by dropping it. The rmcp + /// adapter implements this supervision for its handler futures. + fn execute( + &self, + request: McpRequest, + context: McpRequestContext, + ) -> BoxFuture<'static, Result>; +} diff --git a/src/agent-client-protocol/src/mcp_server/tool_fn.rs b/src/agent-client-protocol/src/mcp_server/tool_fn.rs index 2fa645d4..79a8e569 100644 --- a/src/agent-client-protocol/src/mcp_server/tool_fn.rs +++ b/src/agent-client-protocol/src/mcp_server/tool_fn.rs @@ -1,12 +1,13 @@ //! Runtime-neutral helpers for registering function-backed MCP tools. use futures::{ - SinkExt, StreamExt, - channel::{mpsc, oneshot}, - future::BoxFuture, + StreamExt, + channel::oneshot, + future::{self, BoxFuture, Either}, }; use schemars::JsonSchema; use serde::{Serialize, de::DeserializeOwned}; +use std::pin::Pin; use crate::{ConnectionTo, Error, Role, RunWithConnectionTo}; @@ -16,11 +17,12 @@ struct ToolCall { params: P, mcp_connection: McpConnectionTo, result_tx: futures::channel::oneshot::Sender>, + done_tx: oneshot::Sender<()>, } struct ToolFnMutRunner { func: F, - call_rx: mpsc::Receiver>, + call_rx: Pin>>>, tool_future_fn: Box< dyn for<'a> Fn( &'a mut F, @@ -52,13 +54,34 @@ where while let Some(ToolCall { params, mcp_connection, - result_tx, + mut result_tx, + done_tx, }) = call_rx.next().await { - let result = tool_future_fn(&mut func, params, mcp_connection).await; - result_tx - .send(result) - .map_err(|_| crate::util::internal_error("failed to send MCP result"))?; + // The caller may have cancelled while this invocation waited behind + // another mutable tool call. Do not start work for a gone caller. + if result_tx.is_canceled() { + drop(params); + drop(mcp_connection); + drop(result_tx); + let _ = done_tx.send(()); + continue; + } + let result = { + let cancelled = result_tx.cancellation(); + futures::pin_mut!(cancelled); + match future::select(tool_future_fn(&mut func, params, mcp_connection), cancelled) + .await + { + Either::Left((result, _)) => Some(result), + Either::Right(((), _)) => None, + } + }; + if let Some(result) = result { + // Cancellation after execution is not a runner failure. + drop(result_tx.send(result)); + } + let _ = done_tx.send(()); } Ok(()) } @@ -66,7 +89,7 @@ where struct ToolFnRunner { func: F, - call_rx: mpsc::Receiver>, + call_rx: Pin>>>, tool_future_fn: Box< dyn for<'a> Fn(&'a F, P, McpConnectionTo) -> BoxFuture<'a, Result> + Send @@ -92,9 +115,8 @@ where call_rx, tool_future_fn, } = self; - crate::util::process_stream_concurrently( - call_rx, - async |tool_call| { + call_rx + .for_each_concurrent(64, |tool_call| { fn hack<'a, F, P, R, MyRole>( func: &'a F, params: P, @@ -108,7 +130,8 @@ where + Send + Sync ), - result_tx: oneshot::Sender>, + mut result_tx: oneshot::Sender>, + done_tx: oneshot::Sender<()>, ) -> BoxFuture<'a, ()> where MyRole: Role, @@ -117,8 +140,30 @@ where F: Send + Sync, { Box::pin(async move { - let result = tool_future_fn(func, params, mcp_connection).await; - drop(result_tx.send(result)); + if result_tx.is_canceled() { + drop(params); + drop(mcp_connection); + drop(result_tx); + let _ = done_tx.send(()); + return; + } + let result = { + let cancelled = result_tx.cancellation(); + futures::pin_mut!(cancelled); + match future::select( + tool_future_fn(func, params, mcp_connection), + cancelled, + ) + .await + { + Either::Left((result, _)) => Some(result), + Either::Right(((), _)) => None, + } + }; + if let Some(result) = result { + drop(result_tx.send(result)); + } + let _ = done_tx.send(()); }) } @@ -126,21 +171,27 @@ where params, mcp_connection, result_tx, + done_tx, } = tool_call; - hack(&func, params, mcp_connection, &*tool_future_fn, result_tx).await; - Ok(()) - }, - |a, b| Box::pin(a(b)), - ) - .await + hack( + &func, + params, + mcp_connection, + &*tool_future_fn, + result_tx, + done_tx, + ) + }) + .await; + Ok(()) } } struct ToolFnTool { name: String, description: String, - call_tx: mpsc::Sender>, + call_tx: async_channel::Sender>, } impl McpTool for ToolFnTool @@ -162,13 +213,18 @@ where async fn call_tool(&self, params: P, mcp_connection: McpConnectionTo) -> Result { let (result_tx, result_rx) = oneshot::channel(); + let (done_tx, done_rx) = oneshot::channel(); + #[cfg(feature = "unstable_mcp_over_acp")] + mcp_connection.register_cleanup(done_rx); + #[cfg(not(feature = "unstable_mcp_over_acp"))] + let _done_rx = done_rx; self.call_tx - .clone() .send(ToolCall { params, mcp_connection, result_tx, + done_tx, }) .await .map_err(crate::util::internal_error)?; @@ -192,7 +248,7 @@ pub fn tool_fn_mut( + Send + 'static, ) -> ( - impl McpTool + 'static, + impl McpTool + 'static, impl RunWithConnectionTo, ) where @@ -201,7 +257,7 @@ where Ret: JsonSchema + Serialize + 'static + Send, F: AsyncFnMut(P, McpConnectionTo) -> Result + Send, { - let (call_tx, call_rx) = mpsc::channel(128); + let (call_tx, call_rx) = async_channel::bounded(128); ( ToolFnTool { name: name.to_string(), @@ -210,7 +266,7 @@ where }, ToolFnMutRunner { func, - call_rx, + call_rx: Box::pin(call_rx), tool_future_fn: Box::new(tool_future_fn), }, ) @@ -230,7 +286,7 @@ pub fn tool_fn( + Sync + 'static, ) -> ( - impl McpTool + 'static, + impl McpTool + 'static, impl RunWithConnectionTo, ) where @@ -239,7 +295,7 @@ where Ret: JsonSchema + Serialize + 'static + Send, F: AsyncFn(P, McpConnectionTo) -> Result + Send + Sync + 'static, { - let (call_tx, call_rx) = mpsc::channel(128); + let (call_tx, call_rx) = async_channel::bounded(128); ( ToolFnTool { name: name.to_string(), @@ -248,7 +304,7 @@ where }, ToolFnRunner { func, - call_rx, + call_rx: Box::pin(call_rx), tool_future_fn: Box::new(tool_future_fn), }, ) diff --git a/src/agent-client-protocol/src/role/acp.rs b/src/agent-client-protocol/src/role/acp.rs index f99c1ecc..ba03d6e4 100644 --- a/src/agent-client-protocol/src/role/acp.rs +++ b/src/agent-client-protocol/src/role/acp.rs @@ -886,7 +886,7 @@ fn invalid_initialize_params(error: impl ToString) -> crate::Error { #[cfg(feature = "unstable_protocol_v2")] fn send_initialize_error( - tx: &futures::channel::mpsc::UnboundedSender, + tx: &crate::jsonrpc::FrameSender, frame: &TransportFrame, error: crate::Error, ) -> Result<(), crate::Error> { @@ -944,8 +944,7 @@ fn send_initialize_error( } }; - tx.unbounded_send(response) - .map_err(crate::util::internal_error) + tx.try_send(response).map_err(crate::util::internal_error) } #[cfg(feature = "unstable_protocol_v2")] @@ -972,8 +971,8 @@ async fn reject_initialize( #[cfg(feature = "unstable_protocol_v2")] struct RunningProtocolPeer { - rx: futures::channel::mpsc::UnboundedReceiver, - tx: futures::channel::mpsc::UnboundedSender, + rx: crate::jsonrpc::FrameReceiver, + tx: crate::jsonrpc::FrameSender, future: crate::BoxFuture<'static, Result<(), crate::Error>>, } @@ -981,14 +980,18 @@ struct RunningProtocolPeer { impl RunningProtocolPeer { fn new(component: impl ConnectTo) -> Self { let (Channel { rx, tx }, future) = component.into_channel_and_future(); - Self { rx, tx, future } + Self { + rx, + tx, + future: Box::pin(future), + } } async fn next_frame(self) -> Result, crate::Error> { let Self { mut rx, tx, future } = self; match future::select(Box::pin(rx.next()), future).await { future::Either::Left((Some(frame), future)) => { - Ok(Some((frame, Self { rx, tx, future }))) + Ok(Some((frame.into_frame(), Self { rx, tx, future }))) } future::Either::Left((None, future)) => { future.await?; @@ -1001,7 +1004,7 @@ impl RunningProtocolPeer { return Ok(None); }; Ok(Some(( - frame, + frame.into_frame(), Self { rx, tx, @@ -1024,9 +1027,7 @@ impl RunningProtocolPeer { } fn send_frame(&self, frame: TransportFrame) -> Result<(), crate::Error> { - self.tx - .unbounded_send(frame) - .map_err(crate::util::internal_error) + self.tx.try_send(frame).map_err(crate::util::internal_error) } } diff --git a/src/agent-client-protocol/src/schema/v2_impls.rs b/src/agent-client-protocol/src/schema/v2_impls.rs index 94d849e8..900abbb6 100644 --- a/src/agent-client-protocol/src/schema/v2_impls.rs +++ b/src/agent-client-protocol/src/schema/v2_impls.rs @@ -264,7 +264,27 @@ impl_v2_jsonrpc_request!( ); impl_v2_jsonrpc_request!(v2::PromptRequest, v2::PromptResponse, "session/prompt"); #[cfg(feature = "unstable_mcp_over_acp")] -impl_v2_jsonrpc_request!(v2::MessageMcpRequest, v2::MessageMcpResponse, "mcp/message"); +impl JsonRpcMessage for v2::MessageMcpRequest { + fn matches_method(method: &str) -> bool { + method == "mcp/message" + } + fn method(&self) -> &'static str { + "mcp/message" + } + fn to_untyped_message(&self) -> Result { + UntypedMessage::new("mcp/message", self) + } + fn parse_message(method: &str, params: &impl serde::Serialize) -> Result { + if method != "mcp/message" { + return Err(crate::Error::method_not_found()); + } + crate::util::json_cast_params(params) + } +} +#[cfg(feature = "unstable_mcp_over_acp")] +impl JsonRpcRequest for v2::MessageMcpRequest { + type Response = v2::MessageMcpResponse; +} impl_v2_jsonrpc_notification!(v2::CancelRequestNotification, "$/cancel_request"); impl_v2_jsonrpc_notification!(v2::CancelSessionNotification, "session/cancel"); diff --git a/src/agent-client-protocol/src/util.rs b/src/agent-client-protocol/src/util.rs index ed704bcb..dfc36ef9 100644 --- a/src/agent-client-protocol/src/util.rs +++ b/src/agent-client-protocol/src/util.rs @@ -1,10 +1,5 @@ // Types re-exported from crate root -use futures::{ - future::BoxFuture, - stream::{Stream, StreamExt}, -}; - mod typed; pub use typed::{MatchDispatch, MatchDispatchFrom, TypeNotification}; @@ -113,72 +108,3 @@ pub fn run_until( } }) } - -/// Process items from a stream concurrently. -/// -/// For each item received from `stream`, calls `process_fn` to create a future, -/// then runs all futures concurrently. If any future returns an error, -/// stops processing and returns that error. -/// -/// This is useful for patterns where you receive work items from a channel -/// and want to process them concurrently while respecting backpressure. -pub(crate) async fn process_stream_concurrently( - stream: impl Stream, - process_fn: F, - process_fn_hack: impl for<'a> Fn(&'a F, T) -> BoxFuture<'a, Result<(), crate::Error>>, -) -> Result<(), crate::Error> -where - F: AsyncFn(T) -> Result<(), crate::Error>, -{ - use std::pin::pin; - - use futures::stream::{FusedStream, FuturesUnordered}; - use futures_concurrency::future::Race; - - enum Event { - NewItem(Option), - FutureCompleted(Option>), - } - - let mut stream = pin!(stream.fuse()); - let mut futures: FuturesUnordered<_> = FuturesUnordered::new(); - - loop { - // If we have no futures to run, wait until we do. - if futures.is_empty() { - match stream.next().await { - Some(item) => futures.push(process_fn_hack(&process_fn, item)), - None => return Ok(()), - } - continue; - } - - // If there are no more items coming in, just drain our queue and return. - if stream.is_terminated() { - while let Some(result) = futures.next().await { - result?; - } - return Ok(()); - } - - // Otherwise, race between getting a new item and completing a future. - let event = (async { Event::NewItem(stream.next().await) }, async { - Event::FutureCompleted(futures.next().await) - }) - .race() - .await; - - match event { - Event::NewItem(Some(item)) => { - futures.push(process_fn_hack(&process_fn, item)); - } - Event::FutureCompleted(Some(result)) => { - result?; - } - Event::NewItem(None) | Event::FutureCompleted(None) => { - // Stream closed, loop will catch is_terminated - // No futures were pending, shouldn't happen since we checked is_empty - } - } - } -} diff --git a/src/agent-client-protocol/tests/application_dispatch_v2.rs b/src/agent-client-protocol/tests/application_dispatch_v2.rs index 18943e1f..31267f4d 100644 --- a/src/agent-client-protocol/tests/application_dispatch_v2.rs +++ b/src/agent-client-protocol/tests/application_dispatch_v2.rs @@ -3,8 +3,8 @@ use std::{cell::RefCell, rc::Rc, time::Duration}; use agent_client_protocol::{ - Agent, Channel, Client, Error, RawJsonRpcMessage, TransportBatch, TransportFrame, - V2ConnectionTo, + Agent, BudgetedFrame, Channel, Client, Error, RawJsonRpcMessage, TransportBatch, + TransportFrame, V2ConnectionTo, schema::{ProtocolVersion, v2}, }; use futures::{StreamExt as _, channel::mpsc}; @@ -119,13 +119,13 @@ async fn assert_application_order(batched: bool) { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(initialize))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected initialize"); }; assert_eq!(initialize.method.as_ref(), "initialize"); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( initialize.id, Ok(serde_json::to_value( v2::InitializeResponse::new( @@ -138,9 +138,11 @@ async fn assert_application_order(batched: bool) { ) .unwrap()), ))) + .await .unwrap(); - let Some(TransportFrame::Single(RawJsonRpcMessage::Request(resume))) = peer.rx.next().await + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(resume))) = + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected resume"); }; @@ -168,14 +170,16 @@ async fn assert_application_order(batched: bool) { ]; if batched { peer.tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages(messages).unwrap(), )) + .await .unwrap(); } else { for message in messages { peer.tx - .unbounded_send(TransportFrame::Single(message)) + .send_frame(TransportFrame::Single(message)) + .await .unwrap(); } } diff --git a/src/agent-client-protocol/tests/jsonrpc_advanced.rs b/src/agent-client-protocol/tests/jsonrpc_advanced.rs index 1ed60d53..4f69f6e6 100644 --- a/src/agent-client-protocol/tests/jsonrpc_advanced.rs +++ b/src/agent-client-protocol/tests/jsonrpc_advanced.rs @@ -6,7 +6,7 @@ //! - Out-of-order response handling use agent_client_protocol::{ - Channel, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, JsonRpcMessage, + BudgetedFrame, Channel, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, RawJsonRpcMessage, Responder, SentRequest, TransportBatch, TransportFrame, role::UntypedRole, }; @@ -461,7 +461,7 @@ async fn ordered_callback_installs_dynamic_handler_before_later_batch_entry() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a ping request"); }; @@ -481,7 +481,8 @@ async fn ordered_callback_installs_dynamic_handler_before_later_batch_entry() { let batch = TransportBatch::from_messages([response, notification]) .expect("test batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept the response batch"); while peer.rx.next().await.is_some() {} diff --git a/src/agent-client-protocol/tests/jsonrpc_batch.rs b/src/agent-client-protocol/tests/jsonrpc_batch.rs index 97df192e..8477e3ae 100644 --- a/src/agent-client-protocol/tests/jsonrpc_batch.rs +++ b/src/agent-client-protocol/tests/jsonrpc_batch.rs @@ -865,14 +865,15 @@ async fn protocol_actor_ignores_response_shaped_malformed_public_frame_entries() ]) .expect("test batch is non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("server should accept the test batch"); let frame = tokio::time::timeout(TIMEOUT, peer.rx.next()) .await .expect("timed out waiting for the batch response") .expect("server channel closed before responding"); - let TransportFrame::Batch(batch) = frame else { + let TransportFrame::Batch(batch) = frame.into_frame() else { panic!("request sibling should receive one grouped batch response"); }; let response = serde_json::to_value(batch).expect("batch response should serialize"); diff --git a/src/agent-client-protocol/tests/jsonrpc_error_handling.rs b/src/agent-client-protocol/tests/jsonrpc_error_handling.rs index b1376a36..01fb11e2 100644 --- a/src/agent-client-protocol/tests/jsonrpc_error_handling.rs +++ b/src/agent-client-protocol/tests/jsonrpc_error_handling.rs @@ -90,14 +90,17 @@ async fn response_dispatch_handler_error_reaches_the_local_request_awaiter() { .next() .await .expect("connection should send one request"); - let TransportFrame::Single(RawJsonRpcMessage::Request(request)) = frame else { + let TransportFrame::Single(RawJsonRpcMessage::Request(request)) = + frame.into_frame() + else { panic!("expected one standalone request"); }; peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Ok(serde_json::json!({ "result": "ignored" })), ))) + .await .expect("connection should accept the test response"); Ok::<(), agent_client_protocol::Error>(()) }; diff --git a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs index 4643af1c..b19e929d 100644 --- a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs +++ b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs @@ -12,14 +12,14 @@ use std::{ }; use agent_client_protocol::{ - ByteStreams, Channel, ConnectTo, ConnectionTo, Dispatch, Error, Handled, JsonRpcMessage, - JsonRpcRequest, Lines, RawJsonRpcMessage, TransportFrame, UntypedMessage, + BudgetedFrame, ByteStreams, Channel, ConnectTo, ConnectionTo, Dispatch, Error, Handled, + JsonRpcMessage, JsonRpcRequest, Lines, RawJsonRpcMessage, TransportFrame, UntypedMessage, is_incoming_transport_closed, role::{Role, UntypedRole}, schema::v1::{RequestId, Response}, }; use agent_client_protocol_test::{MyRequest, MyResponse}; -use futures::{FutureExt as _, SinkExt as _, StreamExt as _, future::join, stream}; +use futures::{FutureExt as _, StreamExt as _, future::join, stream}; use tokio::io::{AsyncBufReadExt as _, AsyncWriteExt as _}; use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _}; @@ -90,8 +90,7 @@ impl ConnectTo for PendingTransport { struct QueuedClient { started: futures::channel::oneshot::Sender<()>, - escaped: - futures::channel::oneshot::Sender>, + escaped: futures::channel::oneshot::Sender, } impl ConnectTo for QueuedClient { @@ -103,7 +102,8 @@ impl ConnectTo for QueuedClient { )?; channel .tx - .unbounded_send(TransportFrame::Single(message)) + .send_frame(TransportFrame::Single(message)) + .await .map_err(Error::into_internal_error)?; drop(self.escaped.send(channel.tx.clone())); let _ = self.started.send(()); @@ -180,7 +180,7 @@ fn assert_connection_closed(error: &Error, method: &str) { async fn receive_requests_then_close(mut peer: Channel, count: usize) { for _ in 0..count { assert!(matches!( - peer.rx.next().await, + peer.rx.next().await.map(BudgetedFrame::into_frame), Some(TransportFrame::Single(RawJsonRpcMessage::Request(_))) )); } @@ -188,12 +188,13 @@ async fn receive_requests_then_close(mut peer: Channel, count: usize) { } async fn respond_then_close(mut peer: Channel) { - let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = peer.rx.next().await + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected outgoing request"); }; peer.tx - .send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Ok(serde_json::to_value(MyResponse { status: "received".into(), @@ -326,7 +327,7 @@ async fn channel_peer_receives_final_response_after_write_half_closes() { let peer = async move { let Channel { mut rx, tx } = peer; - tx.unbounded_send(TransportFrame::Single( + tx.send_frame(TransportFrame::Single( RawJsonRpcMessage::request( "myRequest".into(), serde_json::json!({}), @@ -334,6 +335,7 @@ async fn channel_peer_receives_final_response_after_write_half_closes() { ) .unwrap(), )) + .await .expect("channel should accept the final request"); tx.close_channel(); drop(tx); @@ -341,7 +343,7 @@ async fn channel_peer_receives_final_response_after_write_half_closes() { let Some(TransportFrame::Single(RawJsonRpcMessage::Response(Response::Result { id, result, - }))) = rx.next().await + }))) = rx.next().await.map(BudgetedFrame::into_frame) else { panic!("channel read half closed before the final response"); }; @@ -380,7 +382,7 @@ async fn transport_channel_keeps_read_half_open_after_write_half_closes() { let (channel, transport_future) = ConnectTo::::into_channel_and_future(transport); let Channel { mut rx, tx } = channel; - tx.unbounded_send(TransportFrame::Single( + tx.send_frame(TransportFrame::Single( RawJsonRpcMessage::request( "myRequest".into(), serde_json::json!({}), @@ -388,6 +390,7 @@ async fn transport_channel_keeps_read_half_open_after_write_half_closes() { ) .unwrap(), )) + .await .expect("transport channel should accept the request"); tx.close_channel(); drop(tx); @@ -422,7 +425,7 @@ async fn transport_channel_keeps_read_half_open_after_write_half_closes() { let Some(TransportFrame::Single(RawJsonRpcMessage::Response(Response::Result { id, result, - }))) = rx.next().await + }))) = rx.next().await.map(BudgetedFrame::into_frame) else { panic!("read half closed before delivering the peer's final response"); }; @@ -531,7 +534,7 @@ async fn outgoing_drain_keeps_the_full_duplex_read_half_moving() { assert!( escaped - .unbounded_send(TransportFrame::Single( + .try_send(TransportFrame::Single( RawJsonRpcMessage::notification("too-late".into(), serde_json::json!({}),).unwrap() )) .is_err(), @@ -1077,7 +1080,7 @@ async fn request_finishing_conversion_after_eof_keeps_the_eof_cause() { let connection = tokio::spawn(connection); peer.tx - .send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::request( "myRequest".into(), serde_json::json!({}), @@ -1094,7 +1097,8 @@ async fn request_finishing_conversion_after_eof_keeps_the_eof_cause() { assert!(matches!( tokio::time::timeout(TIMEOUT, peer.rx.next()) .await - .expect("handler response was not sent"), + .expect("handler response was not sent") + .map(BudgetedFrame::into_frame), Some(TransportFrame::Single(RawJsonRpcMessage::Response(_))) )); @@ -1233,12 +1237,12 @@ async fn response_buffered_before_eof_is_delivered() { }); let respond_then_close = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected outgoing request"); }; peer.tx - .send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Ok(serde_json::to_value(MyResponse { status: "received".into(), diff --git a/src/agent-client-protocol/tests/protocol_v2.rs b/src/agent-client-protocol/tests/protocol_v2.rs index b6f9579a..562b2d33 100644 --- a/src/agent-client-protocol/tests/protocol_v2.rs +++ b/src/agent-client-protocol/tests/protocol_v2.rs @@ -409,15 +409,16 @@ impl ConnectTo for FutureInitializeV2Client { channel .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "initialize".into(), initialize_params_with_extensions(ProtocolVersion::V2)?, v1::RequestId::Number(1), )?)) + .await .map_err(Error::into_internal_error)?; while let Some(message) = channel.rx.next().await { - let TransportFrame::Single(message) = message else { + let TransportFrame::Single(message) = message.into_frame() else { continue; }; let RawJsonRpcMessage::Response(v1::Response::Result { result, .. }) = message else { @@ -488,15 +489,16 @@ async fn assert_malformed_initialize_rejected(params: Map) -> Res channel .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "initialize".into(), Value::Object(params), v1::RequestId::Number(1), )?)) + .await .map_err(Error::into_internal_error)?; while let Some(message) = channel.rx.next().await { - let TransportFrame::Single(message) = message else { + let TransportFrame::Single(message) = message.into_frame() else { continue; }; let RawJsonRpcMessage::Response(response) = message else { @@ -967,9 +969,8 @@ fn sdk_supported_v2_method_surface_is_jsonrpc_mapped() -> Result<(), Error> { #[cfg(feature = "unstable_mcp_over_acp")] { - fn message_response() -> Result { - serde_json::from_value(serde_json::json!({ "tools": [] })) - .map_err(Error::into_internal_error) + fn message_response() -> v2::MessageMcpResponse { + v2::MessageMcpResponse::success(serde_json::json!({ "tools": [] })) } assert_v2_client_notification_mapping( @@ -988,7 +989,7 @@ fn sdk_supported_v2_method_surface_is_jsonrpc_mapped() -> Result<(), Error> { MessageMcpResponse, "mcp/message", v2::MessageMcpRequest::new("server-1", "request-1", "tools/list"), - message_response()? + message_response() ); } @@ -1047,7 +1048,9 @@ fn mcp_over_acp_v1_variants_are_jsonrpc_mapped() -> Result<(), Error> { assert_response_mapping!( v1::ClientResponse, "mcp/message", - serde_json::json!({ "tools": [] }), + json_value(v1::MessageMcpResponse::success( + serde_json::json!({ "tools": [] }) + ))?, v1::ClientResponse::MessageMcpResponse(_) ); @@ -1953,15 +1956,16 @@ async fn protocol_router_v2_only_rejects_v1_client() -> Result<(), Error> { channel .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "initialize".into(), json_value(v1_initialize_request(ProtocolVersion::V1))?, v1::RequestId::Number(1), )?)) + .await .map_err(Error::into_internal_error)?; while let Some(message) = channel.rx.next().await { - let TransportFrame::Single(message) = message else { + let TransportFrame::Single(message) = message.into_frame() else { continue; }; let RawJsonRpcMessage::Response(v1::Response::Error { error, .. }) = message else { @@ -2699,15 +2703,16 @@ async fn protocol_router_routes_future_protocol_version_to_v2() -> Result<(), Er ); channel .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "initialize".into(), initialize, v1::RequestId::Number(1), )?)) + .await .map_err(Error::into_internal_error)?; while let Some(message) = channel.rx.next().await { - let TransportFrame::Single(message) = message else { + let TransportFrame::Single(message) = message.into_frame() else { continue; }; let RawJsonRpcMessage::Response(v1::Response::Result { result, .. }) = message else { diff --git a/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs b/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs index fab29293..5df47ea8 100644 --- a/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs +++ b/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs @@ -78,18 +78,20 @@ async fn request( let task = tokio::spawn(future); let request_id = v1::RequestId::Number(1); - tx.unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + tx.send_frame(TransportFrame::Single(RawJsonRpcMessage::request( method.into(), params, request_id.clone(), )?)) + .await .map_err(Error::into_internal_error)?; let result = loop { let frame = rx.next().await.ok_or_else(|| { Error::internal_error().data("proxy router closed before initialize response") })?; - let TransportFrame::Single(RawJsonRpcMessage::Response(response)) = frame else { + let TransportFrame::Single(RawJsonRpcMessage::Response(response)) = frame.into_frame() + else { continue; }; match response { diff --git a/src/agent-client-protocol/tests/session_ordering.rs b/src/agent-client-protocol/tests/session_ordering.rs index dee063af..a1850b8c 100644 --- a/src/agent-client-protocol/tests/session_ordering.rs +++ b/src/agent-client-protocol/tests/session_ordering.rs @@ -1,8 +1,8 @@ use std::time::Duration; use agent_client_protocol::{ - ActiveSession, Agent, Channel, Client, Conductor, ConnectionTo, RawJsonRpcMessage, Responder, - SessionMessage, TransportBatch, TransportFrame, + ActiveSession, Agent, BudgetedFrame, Channel, Client, Conductor, ConnectionTo, + RawJsonRpcMessage, Responder, SessionMessage, TransportBatch, TransportFrame, schema::v1::{ ContentBlock, ContentChunk, NewSessionRequest, NewSessionResponse, PromptRequest, PromptResponse, SessionConfigOption, SessionConfigOptionCategory, @@ -34,16 +34,17 @@ async fn initialize_raw_v2_proxy( v2::Implementation::new(client_name, env!("CARGO_PKG_VERSION")), )); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "_proxy/initialize".to_owned(), serde_json::to_value(initialize).expect("initialize request should serialize"), initialize_id.clone(), )?)) + .await .expect("proxy should accept initialization"); let Some(TransportFrame::Single(RawJsonRpcMessage::Response( agent_client_protocol::schema::v1::Response::Result { id, result }, - ))) = peer.rx.next().await + ))) = peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected the proxy initialize response"); }; @@ -306,7 +307,7 @@ async fn on_session_start_installs_routing_before_later_batch_entry() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a session/new request"); }; @@ -333,7 +334,8 @@ async fn on_session_start_installs_routing_before_later_batch_entry() { let batch = TransportBatch::from_messages([response, notification]) .expect("test batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept the response batch"); while peer.rx.next().await.is_some() {} @@ -407,16 +409,17 @@ async fn v2_proxy_session_start_installs_routing_before_later_batch_entry() { let upstream_id = agent_client_protocol::schema::v1::RequestId::Number(2); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "session/new".to_owned(), serde_json::to_value(v2::NewSessionRequest::new("/same-batch-v2-session")) .expect("session request should serialize"), upstream_id.clone(), )?)) + .await .expect("proxy should accept session/new"); let Some(TransportFrame::Single(RawJsonRpcMessage::Request(forwarded))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a forwarded session/new request"); }; @@ -448,13 +451,16 @@ async fn v2_proxy_session_start_installs_routing_before_later_batch_entry() { let batch = TransportBatch::from_messages([response, notification]) .expect("test response batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("proxy should accept the response batch"); let mut saw_response = false; let mut saw_update = false; for _ in 0..2 { - let Some(TransportFrame::Single(message)) = peer.rx.next().await else { + let Some(TransportFrame::Single(message)) = + peer.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected a forwarded session response and update"); }; match message { @@ -558,7 +564,7 @@ async fn v2_proxy_fork_installs_response_id_routing_before_later_batch_entry() { let upstream_id = agent_client_protocol::schema::v1::RequestId::Number(2); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "session/fork".to_owned(), serde_json::to_value(v2::ForkSessionRequest::new( source_session_id, @@ -567,10 +573,11 @@ async fn v2_proxy_fork_installs_response_id_routing_before_later_batch_entry() { .expect("fork request should serialize"), upstream_id.clone(), )?)) + .await .expect("proxy should accept session/fork"); let Some(TransportFrame::Single(RawJsonRpcMessage::Request(forwarded))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a forwarded session/fork request"); }; @@ -603,13 +610,16 @@ async fn v2_proxy_fork_installs_response_id_routing_before_later_batch_entry() { let batch = TransportBatch::from_messages([response, notification]) .expect("test response batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("proxy should accept the response batch"); let mut saw_response = false; let mut saw_update = false; for _ in 0..2 { - let Some(TransportFrame::Single(message)) = peer.rx.next().await else { + let Some(TransportFrame::Single(message)) = + peer.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected a forwarded fork response and update"); }; match message { @@ -714,7 +724,7 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { let upstream_id = agent_client_protocol::schema::v1::RequestId::Number(2); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "session/resume".to_owned(), serde_json::to_value(v2::ResumeSessionRequest::new( session_id.clone(), @@ -723,10 +733,11 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { .expect("resume request should serialize"), upstream_id.clone(), )?)) + .await .expect("proxy should accept session/resume"); let Some(TransportFrame::Single(RawJsonRpcMessage::Request(forwarded))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a forwarded session/resume request"); }; @@ -755,11 +766,12 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { let batch = TransportBatch::from_messages([notification, response]) .expect("test replay batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("proxy should accept the replay batch"); let Some(TransportFrame::Single(RawJsonRpcMessage::Notification(notification))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected replay to be forwarded before the resume response"); }; @@ -775,7 +787,7 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { let Some(TransportFrame::Single(RawJsonRpcMessage::Response( agent_client_protocol::schema::v1::Response::Result { id, result }, - ))) = peer.rx.next().await + ))) = peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected the resume response after replay"); }; diff --git a/src/agent-client-protocol/tests/session_restore.rs b/src/agent-client-protocol/tests/session_restore.rs index c01875e0..6938d5b5 100644 --- a/src/agent-client-protocol/tests/session_restore.rs +++ b/src/agent-client-protocol/tests/session_restore.rs @@ -4,8 +4,9 @@ use std::{future::pending, path::PathBuf, time::Duration}; use agent_client_protocol::{ - Agent, Channel, Client, ConnectionTo, Error, ErrorCode, JsonRpcMessage, JsonRpcNotification, - RawJsonRpcMessage, Responder, SessionMessage, TransportBatch, TransportFrame, UntypedMessage, + Agent, BudgetedFrame, Channel, Client, ConnectionTo, Error, ErrorCode, JsonRpcMessage, + JsonRpcNotification, RawJsonRpcMessage, Responder, SessionMessage, TransportBatch, + TransportFrame, UntypedMessage, schema::v1::{ CancelRequestNotification, ContentBlock, ContentChunk, LoadSessionRequest, LoadSessionResponse, RequestId, ResumeSessionRequest, ResumeSessionResponse, @@ -138,7 +139,7 @@ async fn load_session_preserves_pre_response_replay_and_exact_response() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/load") }; @@ -153,7 +154,8 @@ async fn load_session_preserves_pre_response_replay_and_exact_response() { ]) .expect("restore batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept replay and response"); while peer.rx.next().await.is_some() {} @@ -206,7 +208,7 @@ async fn resume_session_returns_exact_response_and_an_active_session() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/resume") }; @@ -219,7 +221,8 @@ async fn resume_session_returns_exact_response_and_an_active_session() { ]) .expect("resume batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept update and response"); while peer.rx.next().await.is_some() {} @@ -256,7 +259,7 @@ async fn resume_session_from_preserves_the_existing_request() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/resume") }; @@ -265,10 +268,11 @@ async fn resume_session_from_preserves_the_existing_request() { peer_request ); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Ok(serde_json::to_value(ResumeSessionResponse::new())?), ))) + .await .expect("client should accept resume response"); while peer.rx.next().await.is_some() {} Ok::<(), Error>(()) @@ -314,11 +318,12 @@ async fn restore_waits_for_routing_acknowledgment_before_publication() { let peer = async move { peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "test/block-incoming".to_owned(), serde_json::json!({}), RequestId::Number(1), )?)) + .await .expect("client should accept the blocking request"); restore_called_rx .await @@ -381,7 +386,7 @@ async fn failed_restore_removes_routing_before_later_batch_entries() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/load") }; @@ -395,18 +400,21 @@ async fn failed_restore_removes_routing_before_later_batch_entries() { ]) .expect("failure batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept failure, probe, and barrier"); - let Some(TransportFrame::Single(RawJsonRpcMessage::Request(retry))) = peer.rx.next().await + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(retry))) = + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected retry session/load") }; peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( retry.id, Ok(serde_json::to_value(LoadSessionResponse::new())?), ))) + .await .expect("client should accept retry response"); while peer.rx.next().await.is_some() {} Ok::<(), Error>(()) @@ -477,7 +485,7 @@ async fn cancelling_restore_cancels_request_and_removes_routing() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/resume") }; @@ -490,7 +498,7 @@ async fn cancelling_restore_cancels_request_and_removes_routing() { .map_err(Error::into_internal_error)?; let Some(TransportFrame::Single(RawJsonRpcMessage::Notification(notification))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("dropping the restore future should send $/cancel_request") }; @@ -504,28 +512,32 @@ async fn cancelling_restore_cancels_request_and_removes_routing() { ]) .expect("cancellation probe batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(probe_batch)) + .send_frame(TransportFrame::Batch(probe_batch)) + .await .expect("client should accept cancellation probe and barrier"); barrier_observed_rx .await .map_err(Error::into_internal_error)?; peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Err(Error::request_cancelled()), ))) + .await .expect("client should accept the cancelled request's response"); - let Some(TransportFrame::Single(RawJsonRpcMessage::Request(retry))) = peer.rx.next().await + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(retry))) = + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected retry session/resume") }; peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( retry.id, Ok(serde_json::to_value(ResumeSessionResponse::new())?), ))) + .await .expect("client should accept retry response"); while peer.rx.next().await.is_some() {} Ok::<(), Error>(()) diff --git a/src/agent-client-protocol/tests/session_v2_mcp.rs b/src/agent-client-protocol/tests/session_v2_mcp.rs index f9ea571a..81009032 100644 --- a/src/agent-client-protocol/tests/session_v2_mcp.rs +++ b/src/agent-client-protocol/tests/session_v2_mcp.rs @@ -211,7 +211,13 @@ async fn run_mcp_round_trip( ) .block_task() .await?; - let response = serde_json::from_str(response.0.get()).map_err(Error::into_internal_error)?; + let response = match response { + v2::MessageMcpResponse::Result { result, .. } => result, + v2::MessageMcpResponse::Error { error, .. } => { + return Err(Error::new(error.code, error.message)); + } + _ => return Err(Error::internal_error().data("unknown MCP response carrier")), + }; Ok(RoundTrip { server_id: server_id.to_string(), From 3de30faf8c8beaf8fe3bfb2ae85efb9b46e37eaf Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Fri, 25 Sep 2026 13:18:18 +0200 Subject: [PATCH 4/8] refactor(acp): adapt independent versioned MCP outcomes --- Cargo.lock | 2 +- Cargo.toml | 2 +- md/migration-stateless-mcp.md | 6 +- md/protocol.md | 2 +- .../src/mcp_over_acp/mod.rs | 30 +++-- .../src/mcp_over_acp/protocol.rs | 32 ++++- .../src/mcp_server/active_session.rs | 115 ++++++++++++++---- .../src/schema/v2_impls.rs | 22 +--- 8 files changed, 149 insertions(+), 62 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index ed05ae40..31b26d02 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -141,7 +141,7 @@ dependencies = [ [[package]] name = "agent-client-protocol-schema" version = "1.9.1" -source = "git+https://github.com/agentclientprotocol/agent-client-protocol?rev=9d8b499a332ffa0655007d4be7dd949e05180de3#9d8b499a332ffa0655007d4be7dd949e05180de3" +source = "git+https://github.com/agentclientprotocol/agent-client-protocol?rev=e5c36d2671fd355f983533bc83b5feb7981d25a6#e5c36d2671fd355f983533bc83b5feb7981d25a6" dependencies = [ "anyhow", "derive_more", diff --git a/Cargo.toml b/Cargo.toml index efa72e48..db895a06 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -36,7 +36,7 @@ yopo = { package = "agent-client-protocol-yopo", path = "src/yopo" } # Protocol # Draft cross-repository validation; replace with the released schema before publishing. -agent-client-protocol-schema = { git = "https://github.com/agentclientprotocol/agent-client-protocol", rev = "9d8b499a332ffa0655007d4be7dd949e05180de3", default-features = false, features = ["tracing"] } +agent-client-protocol-schema = { git = "https://github.com/agentclientprotocol/agent-client-protocol", rev = "e5c36d2671fd355f983533bc83b5feb7981d25a6", default-features = false, features = ["tracing"] } # Core async runtime tokio = { version = "1.52", default-features = false } diff --git a/md/migration-stateless-mcp.md b/md/migration-stateless-mcp.md index 6b26975f..39f03d31 100644 --- a/md/migration-stateless-mcp.md +++ b/md/migration-stateless-mcp.md @@ -53,9 +53,9 @@ Do not run ACP authentication handling on an inner MCP error code. A tool execution failure with `isError` remains an MCP result. MRTR's `input_required` also remains a result; retry with fresh IDs/metadata and unchanged opaque state. -Both ACP versions export the same response/error carrier types. Downstream -code that implements traits for these types must not provide separate v1 and -v2 implementations. +ACP v1 and v2 define independent response/error carrier types. They currently +use the same JSON representation, but may evolve separately. Use the types +for the negotiated ACP version and keep trait implementations version-specific. ## Separate services from operations diff --git a/md/protocol.md b/md/protocol.md index 9ff0a2f3..6528391c 100644 --- a/md/protocol.md +++ b/md/protocol.md @@ -125,7 +125,7 @@ The successful outer ACP response contains exactly one MCP outcome: ``` An MCP protocol error uses `{"error": {"code": ..., "message": ..., "data": ...}}` -inside the successful outer `result`, not an ACP error response. The shared +inside the successful outer `result`, not an ACP error response. Each version's `MessageMcpResponse::{Result, Error}` type preserves this distinction. Inner results are opaque JSON (including null); inner error data distinguishes null from omission. MCP error codes never acquire ACP meanings. diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs index 0b0262ee..26689737 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs @@ -16,7 +16,7 @@ use std::{ use agent_client_protocol::{ Agent, Client, Conductor, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - Proxy, UntypedMessage, schema::v1::MessageMcpResponse, util::MatchDispatchFrom, + Proxy, UntypedMessage, util::MatchDispatchFrom, }; use futures::{ SinkExt, StreamExt, @@ -26,7 +26,9 @@ use serde_json::Value; use tokio::{net::TcpListener, sync::mpsc as tokio_mpsc}; use tracing::{debug, warn}; -use self::protocol::{DownstreamMcpMode, NativeMcpNotification, PolyfillProtocol}; +use self::protocol::{ + DownstreamMcpMode, NativeMcpNotification, NativeMcpOutcome, PolyfillProtocol, +}; // Conservative per-bridge limits. Notifications are bounded per HTTP POST by // both message count and serialized bytes; terminal responses bypass the queue. @@ -327,6 +329,7 @@ async fn transform_session_servers( } struct ActiveRequest { + protocol: PolyfillProtocol, server_id: String, http_id: Value, method: String, @@ -408,6 +411,7 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { self.active.insert( request_id.clone(), ActiveRequest { + protocol, server_id: server_id.clone(), http_id: http_id.clone(), method: method.clone(), @@ -480,6 +484,7 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { let http_id = active.http_id.clone(); let value = match result { Ok(carrier) => project_mcp_carrier( + active.protocol, active.http_id, &request_id, &active.method, @@ -509,20 +514,21 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { /// ACP success carries exactly one MCP outcome. An outer ACP failure is a /// binding/runtime failure, not an MCP error carried in a successful response. -fn project_mcp_carrier(http_id: Value, request_id: &str, method: &str, carrier: Value) -> Value { - // Both ACP revisions share this type. Keep envelope validation in the schema, - // rather than maintaining a second parser that can drift from its null rules. - match serde_json::from_value::(carrier) { - Ok(MessageMcpResponse::Result { mut result, .. }) => { +fn project_mcp_carrier( + protocol: PolyfillProtocol, + http_id: Value, + request_id: &str, + method: &str, + carrier: Value, +) -> Value { + match protocol.message_response(carrier) { + Ok(NativeMcpOutcome::Result(mut result)) => { if method == "tools/list" { strip_header_annotations(&mut result); } http::rpc_result(http_id, request_id, result) } - Ok(MessageMcpResponse::Error { error, .. }) => http::rpc_peer_error( - http_id, - serde_json::to_value(error).expect("MCP errors contain only JSON values"), - ), + Ok(NativeMcpOutcome::Error(error)) => http::rpc_peer_error(http_id, error), _ => http::rpc_error(http_id, -33002, "Invalid MCP-over-ACP response carrier"), } } @@ -753,6 +759,7 @@ mod http_limits_tests { runner.active.insert( index.to_string(), ActiveRequest { + protocol: PolyfillProtocol::V1, server_id: String::new(), http_id: Value::Null, method: String::new(), @@ -780,6 +787,7 @@ mod tests { }); let project = |carrier| { project_mcp_carrier( + PolyfillProtocol::V1, serde_json::json!("external"), "internal", "tools/call", diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs index d6d7aea5..ea595f4a 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs @@ -2,7 +2,7 @@ use agent_client_protocol::{ Error, JsonRpcMessage, JsonRpcResponse, UntypedMessage, schema::{ InitializeProxyRequest, METHOD_INITIALIZE_PROXY, ProtocolVersion, - v1::{LoadSessionRequest, McpServer, NewSessionRequest, ResumeSessionRequest}, + v1::{self, LoadSessionRequest, McpServer, NewSessionRequest, ResumeSessionRequest}, }, }; use serde_json::{Map, Value}; @@ -19,7 +19,37 @@ pub(crate) enum PolyfillProtocol { V2, } +pub(super) enum NativeMcpOutcome { + Result(Value), + Error(Value), +} + impl PolyfillProtocol { + /// Validate against the negotiated ACP version before projecting onto HTTP. + pub(super) fn message_response(self, value: Value) -> Result { + match self { + Self::V1 => match v1::MessageMcpResponse::from_value("mcp/message", value)? { + v1::MessageMcpResponse::Result { result, .. } => { + Ok(NativeMcpOutcome::Result(result)) + } + v1::MessageMcpResponse::Error { error, .. } => { + Ok(NativeMcpOutcome::Error(serde_json::to_value(error)?)) + } + _ => Err(Error::invalid_request().data("unsupported MCP outcome")), + }, + #[cfg(feature = "unstable_protocol_v2")] + Self::V2 => match v2::MessageMcpResponse::from_value("mcp/message", value)? { + v2::MessageMcpResponse::Result { result, .. } => { + Ok(NativeMcpOutcome::Result(result)) + } + v2::MessageMcpResponse::Error { error, .. } => { + Ok(NativeMcpOutcome::Error(serde_json::to_value(error)?)) + } + _ => Err(Error::invalid_request().data("unsupported MCP outcome")), + }, + } + } + pub(crate) fn from_initialize_request(request: &UntypedMessage) -> Result { if request.method() != METHOD_INITIALIZE_PROXY { return Err(Error::invalid_request().data("expected initialize proxy request")); diff --git a/src/agent-client-protocol/src/mcp_server/active_session.rs b/src/agent-client-protocol/src/mcp_server/active_session.rs index 0d3ed260..1df3fc3d 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -15,7 +15,8 @@ use std::{ use crate::{ Agent, Channel, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - JsonRpcNotification, JsonRpcRequest, RawJsonRpcMessage, Responder, Role, TransportFrame, + JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, RawJsonRpcMessage, Responder, Role, + TransportFrame, mcp_server::{ MCP_BACKEND_FAILURE, MCP_RESOURCE_EXHAUSTED, McpConnectionContext, McpConnectionTo, McpOperationCancellation, McpOutcome, McpRequest, McpRequestContext, McpServerConnect, @@ -40,9 +41,11 @@ pub(super) struct V1McpProtocol; pub(super) struct V2McpProtocol; pub(super) trait McpProtocol: Send + 'static { - type MessageRequest: JsonRpcRequest; + type MessageRequest: JsonRpcRequest; + type MessageResponse: JsonRpcResponse + serde::Serialize; type MessageNotification: JsonRpcNotification; + fn response(outcome: McpOutcome) -> Self::MessageResponse; fn server_id(request: &Self::MessageRequest) -> McpServerAcpId; fn request_id(request: &Self::MessageRequest) -> McpRequestId; fn into_request(request: Self::MessageRequest) -> (String, Option>); @@ -56,8 +59,16 @@ pub(super) trait McpProtocol: Send + 'static { impl McpProtocol for V1McpProtocol { type MessageRequest = MessageMcpRequest; + type MessageResponse = MessageMcpResponse; type MessageNotification = MessageMcpNotification; + fn response(outcome: McpOutcome) -> Self::MessageResponse { + match outcome { + McpOutcome::Result(value) => MessageMcpResponse::success(value), + McpOutcome::Error(error) => MessageMcpResponse::error(error), + } + } + fn server_id(request: &Self::MessageRequest) -> McpServerAcpId { request.server_id.clone() } @@ -85,11 +96,10 @@ fn into_mcp_error(error: crate::Error) -> McpError { mcp } -fn outcome_response(outcome: McpOutcome) -> Result { - let response = match outcome { - McpOutcome::Result(value) => MessageMcpResponse::success(value), - McpOutcome::Error(error) => MessageMcpResponse::error(error), - }; +fn outcome_response( + outcome: McpOutcome, +) -> Result { + let response = Protocol::response(outcome); check_payload_size(&response, MAX_PAYLOAD_BYTES)?; Ok(response) } @@ -104,8 +114,8 @@ fn project_outcome(outcome: McpOutcome, is_discovery: bool) -> Result, +fn send_outcome( + responder: Responder, result: Result, is_discovery: bool, ) -> Result<(), crate::Error> { @@ -115,7 +125,7 @@ fn send_outcome( // remain named outer ACP errors, regardless of backend type. let outcome = project_outcome(outcome, is_discovery) .unwrap_or_else(|error| McpOutcome::Error(into_mcp_error(error))); - match outcome_response(outcome) { + match outcome_response::(outcome) { Ok(response) => responder.respond(response), Err(error) => responder.respond_with_error(error), } @@ -127,8 +137,23 @@ fn send_outcome( #[cfg(feature = "unstable_protocol_v2")] impl McpProtocol for V2McpProtocol { type MessageRequest = crate::schema::v2::MessageMcpRequest; + type MessageResponse = crate::schema::v2::MessageMcpResponse; type MessageNotification = crate::schema::v2::MessageMcpNotification; + fn response(outcome: McpOutcome) -> Self::MessageResponse { + match outcome { + McpOutcome::Result(value) => Self::MessageResponse::success(value), + McpOutcome::Error(error) => { + // The service outcome uses the v1 error representation. Adapt it + // explicitly here instead of coupling the versioned wire types. + let mut wire_error = crate::schema::v2::McpError::new(error.code, error.message); + wire_error.data = error.data; + wire_error.extra = error.extra; + Self::MessageResponse::error(wire_error) + } + } + } + fn server_id(request: &Self::MessageRequest) -> McpServerAcpId { McpServerAcpId::new(request.server_id.0.clone()) } @@ -242,10 +267,15 @@ where fn handle_request( &mut self, request: Protocol::MessageRequest, - responder: Responder, + responder: Responder, connection: &ConnectionTo, - ) -> Result)>, crate::Error> - { + ) -> Result< + Handled<( + Protocol::MessageRequest, + Responder, + )>, + crate::Error, + > { let server_id = Protocol::server_id(&request); if server_id != self.server_id { return Ok(Handled::No { @@ -261,7 +291,9 @@ where return Ok(Handled::Yes); } if let Err(error) = validate_modern_request(&method, params.as_ref()) { - responder.respond(outcome_response(McpOutcome::Error(into_mcp_error(error)))?)?; + responder.respond(outcome_response::(McpOutcome::Error( + into_mcp_error(error), + ))?)?; return Ok(Handled::Yes); } let (guard, stop_rx) = match admit_request(&self.active, request_id.clone()) { @@ -376,7 +408,7 @@ where cleanup_connection.wait_cleanup().await; // Operation futures have been dropped and cannot send late output. drop(guard); - let response = send_outcome(responder, result, is_discovery); + let response = send_outcome::(responder, result, is_discovery); if let Err(error) = response { tracing::debug!(?error, "cannot send MCP response"); } @@ -533,7 +565,7 @@ where drop(backend_done_rx.await); cleanup_connection.wait_cleanup().await; drop(guard); - let response = send_outcome(responder, result, is_discovery); + let response = send_outcome::(responder, result, is_discovery); if let Err(error) = response { tracing::debug!(?error, "cannot send request-scoped MCP response"); } @@ -629,9 +661,9 @@ fn validate_modern_request( #[cfg(test)] mod tests { use super::{ - ActiveRequests, MAX_ACTIVE_REQUESTS, MAX_PAYLOAD_BYTES, McpOutcome, admit_request, - check_payload_size, constrain_discovery_versions, into_mcp_error, outcome_response, - validate_modern_request, + ActiveRequests, MAX_ACTIVE_REQUESTS, MAX_PAYLOAD_BYTES, McpOutcome, V1McpProtocol, + admit_request, check_payload_size, constrain_discovery_versions, into_mcp_error, + outcome_response, validate_modern_request, }; use crate::{ mcp_server::MCP_RESOURCE_EXHAUSTED, @@ -647,16 +679,51 @@ mod tests { McpError::new(-32000, "peer error").data(json!("x".repeat(MAX_PAYLOAD_BYTES))), ), ] { - let error = outcome_response(outcome).expect_err("oversized carrier must be rejected"); + let error = outcome_response::(outcome) + .expect_err("oversized carrier must be rejected"); assert_eq!(i32::from(error.code), MCP_RESOURCE_EXHAUSTED); } - let result = outcome_response(McpOutcome::Result(serde_json::Value::Null)).unwrap(); + let result = + outcome_response::(McpOutcome::Result(serde_json::Value::Null)).unwrap(); assert_eq!( serde_json::to_value(result).unwrap(), json!({"result":null}) ); } + #[cfg(feature = "unstable_protocol_v2")] + #[test] + fn versioned_outcomes_preserve_results_and_error_fields() { + for value in [ + serde_json::Value::Null, + json!({"resultType":"complete","_meta":{"custom":true}}), + ] { + let v1 = outcome_response::(McpOutcome::Result(value.clone())).unwrap(); + let v2 = outcome_response::(McpOutcome::Result(value)).unwrap(); + assert_eq!( + serde_json::to_value(v1).unwrap(), + serde_json::to_value(v2).unwrap() + ); + } + for data in [ + None, + Some(serde_json::Value::Null), + Some(json!({"details":[1,2]})), + ] { + let mut error = McpError::new(-32000, "opaque peer error"); + if let Some(data) = data { + error = error.data(data); + } + error + .extra + .insert("extension".into(), json!({"preserve":true})); + let expected = json!({"error": error}); + let v2: crate::schema::v2::MessageMcpResponse = + outcome_response::(McpOutcome::Error(error)).unwrap(); + assert_eq!(serde_json::to_value(v2).unwrap(), expected); + } + } + #[test] fn only_modern_request_metadata_is_accepted() { let modern = json!({"_meta": {"io.modelcontextprotocol/protocolVersion": "2026-07-28", "io.modelcontextprotocol/clientCapabilities": {}, "requestState": {"opaque": true}}}); @@ -753,12 +820,14 @@ mod tests { .as_object(), ) .expect_err("unsupported inner version"); - let response = outcome_response(McpOutcome::Error(into_mcp_error(unsupported))).unwrap(); + let response = + outcome_response::(McpOutcome::Error(into_mcp_error(unsupported))) + .unwrap(); assert_eq!( serde_json::to_value(response).unwrap()["error"]["code"], -32022 ); - let response = outcome_response(McpOutcome::Error( + let response = outcome_response::(McpOutcome::Error( McpError::new(-32000, "opaque MCP error").data(serde_json::Value::Null), )) .unwrap(); diff --git a/src/agent-client-protocol/src/schema/v2_impls.rs b/src/agent-client-protocol/src/schema/v2_impls.rs index 900abbb6..94d849e8 100644 --- a/src/agent-client-protocol/src/schema/v2_impls.rs +++ b/src/agent-client-protocol/src/schema/v2_impls.rs @@ -264,27 +264,7 @@ impl_v2_jsonrpc_request!( ); impl_v2_jsonrpc_request!(v2::PromptRequest, v2::PromptResponse, "session/prompt"); #[cfg(feature = "unstable_mcp_over_acp")] -impl JsonRpcMessage for v2::MessageMcpRequest { - fn matches_method(method: &str) -> bool { - method == "mcp/message" - } - fn method(&self) -> &'static str { - "mcp/message" - } - fn to_untyped_message(&self) -> Result { - UntypedMessage::new("mcp/message", self) - } - fn parse_message(method: &str, params: &impl serde::Serialize) -> Result { - if method != "mcp/message" { - return Err(crate::Error::method_not_found()); - } - crate::util::json_cast_params(params) - } -} -#[cfg(feature = "unstable_mcp_over_acp")] -impl JsonRpcRequest for v2::MessageMcpRequest { - type Response = v2::MessageMcpResponse; -} +impl_v2_jsonrpc_request!(v2::MessageMcpRequest, v2::MessageMcpResponse, "mcp/message"); impl_v2_jsonrpc_notification!(v2::CancelRequestNotification, "$/cancel_request"); impl_v2_jsonrpc_notification!(v2::CancelSessionNotification, "session/cancel"); From 32f0c4fe34f84185a058612d21b9058ffc828804 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Fri, 25 Sep 2026 13:34:43 +0200 Subject: [PATCH 5/8] test(acp): verify MCP request and notification decoding Cover outer-id discrimination and separate v1/v2 method enums, and clarify that the binding adds no stronger cancellation support requirement. --- md/protocol.md | 7 +- .../tests/mcp_message_deserialization.rs | 73 +++++++++++++++++++ 2 files changed, 77 insertions(+), 3 deletions(-) create mode 100644 src/agent-client-protocol/tests/mcp_message_deserialization.rs diff --git a/md/protocol.md b/md/protocol.md index 6528391c..d342846b 100644 --- a/md/protocol.md +++ b/md/protocol.md @@ -169,12 +169,13 @@ mean no parameters. A valid modern request still needs its required Use [`$/cancel_request`](./request-cancellation.md) with the outer ACP request ID. Normal proxy forwarding maps this cancellation hop by hop. It never -rewrites the logical MCP ID. Advertising this binding requires cancellation -handling even where the underlying ACP revision makes general cancellation optional. +rewrites the logical MCP ID. Cancellation is best effort; advertising this +transport does not guarantee that every operation can be cancelled or impose +an additional cancellation support requirement. Each operation owns its backend work. A result, error, cancellation, or registration removal ends that operation; sibling requests and subscriptions stay -independent. Cancellation revokes output immediately, but the operation keeps its +independent. When the SDK honors cancellation, it revokes output but keeps the admission slot and logical ID until owned cleanup finishes. There is no MCP connection ID to release. `server/discover` is an ordinary optional request, not a prerequisite for tool calls. diff --git a/src/agent-client-protocol/tests/mcp_message_deserialization.rs b/src/agent-client-protocol/tests/mcp_message_deserialization.rs new file mode 100644 index 00000000..dd8c3fe0 --- /dev/null +++ b/src/agent-client-protocol/tests/mcp_message_deserialization.rs @@ -0,0 +1,73 @@ +//! A shared method name must not conflate JSON-RPC requests and notifications. +#![cfg(feature = "unstable_mcp_over_acp")] + +use agent_client_protocol::{JsonRpcMessage, RawJsonRpcMessage, schema::v1}; +use serde_json::{Value, json}; + +fn params() -> Value { + // Deliberately identical for both message kinds: dispatch must use the outer + // envelope, not infer a kind from this opaque inner method or requestId. + json!({"serverId":"server", "requestId":"logical-id", "method":"custom/message"}) +} + +#[test] +fn mcp_message_kind_is_selected_by_outer_id() { + let notification = json!({"jsonrpc":"2.0", "method":"mcp/message", "params":params()}); + let parsed: RawJsonRpcMessage = serde_json::from_value(notification.clone()).unwrap(); + let RawJsonRpcMessage::Notification(parsed) = parsed else { + panic!("nested requestId must not turn a notification into a request"); + }; + assert_eq!(parsed.method.as_ref(), "mcp/message"); + assert_eq!(parsed.params.unwrap().into_value(), params()); + + for id in [json!(42), json!("outer-id")] { + let mut request = notification.clone(); + request["id"] = id.clone(); + let parsed: RawJsonRpcMessage = serde_json::from_value(request).unwrap(); + let RawJsonRpcMessage::Request(parsed) = parsed else { + panic!("outer id identifies a request"); + }; + assert_eq!(serde_json::to_value(parsed.id).unwrap(), id); + assert_eq!(parsed.method.as_ref(), "mcp/message"); + assert_eq!(parsed.params.unwrap().into_value(), params()); + } + + for id in [json!(true), json!({}), json!([])] { + let mut malformed = notification.clone(); + malformed["id"] = id; + assert!( + serde_json::from_value::(malformed).is_err(), + "an invalid request id must not fall back to notification deserialization" + ); + } +} + +#[test] +fn v1_mcp_method_is_in_separate_request_and_notification_enums() { + assert!(matches!( + v1::AgentRequest::parse_message("mcp/message", ¶ms()).unwrap(), + v1::AgentRequest::MessageMcpRequest(_) + )); + assert!(matches!( + v1::ClientNotification::parse_message("mcp/message", ¶ms()).unwrap(), + v1::ClientNotification::MessageMcpNotification(_) + )); + assert!(v1::ClientRequest::parse_message("mcp/message", ¶ms()).is_err()); + assert!(v1::AgentNotification::parse_message("mcp/message", ¶ms()).is_err()); +} + +#[cfg(feature = "unstable_protocol_v2")] +#[test] +fn v2_mcp_method_is_in_separate_request_and_notification_enums() { + use agent_client_protocol::schema::v2; + assert!(matches!( + v2::AgentRequest::parse_message("mcp/message", ¶ms()).unwrap(), + v2::AgentRequest::MessageMcpRequest(_) + )); + assert!(matches!( + v2::ClientNotification::parse_message("mcp/message", ¶ms()).unwrap(), + v2::ClientNotification::MessageMcpNotification(_) + )); + assert!(v2::ClientRequest::parse_message("mcp/message", ¶ms()).is_err()); + assert!(v2::AgentNotification::parse_message("mcp/message", ¶ms()).is_err()); +} From f54f7d078a08128ff082704617c86bbb22b5e5e9 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Mon, 28 Sep 2026 12:32:42 +0200 Subject: [PATCH 6/8] fix(runtime): preserve bounded progress and shutdown cleanup Keep cancellation observable through readiness gates, bound all live tasks and frame sink reservations, and join protected MCP cleanup through connection termination. Repair cold HTTP session setup, release consumed output permits, abort rejected WebSocket input, and retain measured session metadata. Add regressions for the review findings and fix the minimal-feature build. --- md/http-transport.md | 24 +- md/mcp-over-acp.md | 11 +- md/request-cancellation.md | 15 + md/transport-architecture.md | 66 +- .../tests/initialization_v2.rs | 24 +- .../tests/request_cancellation.rs | 74 +- src/agent-client-protocol-http/src/client.rs | 136 ++-- .../src/connection.rs | 39 +- .../src/connection_admission_tests.rs | 108 ++- src/agent-client-protocol-http/src/server.rs | 181 +++++ .../src/websocket_server.rs | 177 ++++- .../src/mcp_over_acp/http.rs | 52 +- .../tests/stateless_native_mcp.rs | 107 ++- src/agent-client-protocol/src/jsonrpc.rs | 662 +++++++++++++++--- .../src/jsonrpc/admission.rs | 22 +- src/agent-client-protocol/src/jsonrpc/run.rs | 27 +- .../src/jsonrpc/task_actor.rs | 127 +++- .../src/mcp_server/active_session.rs | 6 +- .../src/mcp_server/context.rs | 2 +- .../tests/jsonrpc_request_cancellation.rs | 47 +- .../tests/live_task_capacity.rs | 80 +++ .../tests/native_mcp_shutdown.rs | 220 ++++++ 22 files changed, 1949 insertions(+), 258 deletions(-) create mode 100644 src/agent-client-protocol/tests/live_task_capacity.rs create mode 100644 src/agent-client-protocol/tests/native_mcp_shutdown.rs diff --git a/md/http-transport.md b/md/http-transport.md index c8584e76..8f97b804 100644 --- a/md/http-transport.md +++ b/md/http-transport.md @@ -59,10 +59,26 @@ active session, clients should also open: - `Acp-Connection-Id: ` - `Acp-Session-Id: ` -Open a session stream before sending methods such as `session/prompt`, -`session/load`, `session/resume`, or other session-scoped requests. When a -`session/new` or `session/fork` response returns a new `sessionId`, open an SSE -stream for that returned session before expecting updates or responses for it. +For a session not yet used on this HTTP connection, send its first +session-scoped POST (for example, `session/load` or `session/resume`) and wait +for `202 Accepted` before opening its GET. The POST registers the bounded +session mailbox before it is admitted to the agent. History and other output +can then queue until the GET attaches; an unknown-session GET returns 409 and +does not allocate a mailbox. A batch registers all its session mailboxes before +returning 202. + +Keep the connection stream and existing session streams running during this +setup so callback responses and cancellation can still progress. `HttpClient` +does this automatically. When a `session/new` or `session/fork` response returns +a new `sessionId`, its mailbox is already registered and the client can open +the corresponding GET immediately. + +Frame and pending-work limits apply to HTTP and WebSocket traffic. Session +metadata retains a measured charge for its ID, not the entire opening request. +Consuming an SSE/WebSocket frame releases its payload charge before waiting +for the next frame. A WebSocket frame rejected by admission terminates that +connection rather than silently losing input or waiting forever to drain a +still-live agent. ## Features diff --git a/md/mcp-over-acp.md b/md/mcp-over-acp.md index 09e479da..0dd66971 100644 --- a/md/mcp-over-acp.md +++ b/md/mcp-over-acp.md @@ -34,7 +34,10 @@ cancellation and cleanup instead of merely dropping detached task handles. Custom `McpService` implementations must observe `operation_cancellation()` and return only after their owned cleanup finishes. The binding waits for this -completion; it cannot forcibly terminate detached application work. +completion, including on ACP EOF, a runtime error, or a `connect_with` +foreground return, before dropping operation supervisors or scoped tool +runners. It does not join arbitrary user-spawned tasks or forcibly terminate +detached application work. The scoped `tool_fn` helpers continue to provide `McpConnectionTo` for host ACP access. For decisions using the full MCP metadata/capabilities, implement @@ -101,9 +104,13 @@ ownership. Adapters must keep the frame's permit through staging, deferred dispatch, and writes; extracting a payload must not silently release its charge while retaining the data. Async producers await capacity; synchronous dispatch must fail explicitly instead of blocking the dispatcher needed to free capacity. +`max_queued_frames` bounds all live SDK tasks (running plus waiting), not just +waiting task slots. A persistent child connection may occupy one slot; when no +live slot remains, an ordered response callback is rejected immediately rather +than accepted behind a child that cannot finish. The same item-limit policy currently governs frame queues, pending requests, -running tasks, dynamic handlers, and deferred dispatch; the default is 32. +live tasks, dynamic handlers, and deferred dispatch; the default is 32. The shared payload budget defaults to 64 MiB with a 16 MiB frame maximum and reserved response/cancellation capacity. These are serialized-payload charges, not an exact bound on total process memory or allocations inside user code. diff --git a/md/request-cancellation.md b/md/request-cancellation.md index 5c3a56b4..681cc162 100644 --- a/md/request-cancellation.md +++ b/md/request-cancellation.md @@ -42,6 +42,21 @@ an unknown or already-completed request ID is silently ignored. A not a string, number, or null) is logged and ignored without a reply, like any other malformed notification. +### Cancelling before publication + +`send_request` queues a request; it does not prove that the peer has received it. +If cancellation reaches the outgoing actor before it publishes that request, +the SDK settles it locally with `-32800` and sends neither the request nor its +cancellation notification. This also applies to requests waiting for session +readiness. Once published, the peer's cooperative cancellation rules apply. + +Cancellation uses a separate bounded urgent queue so it can bypass a blocked +readiness gate regardless of ordinary queue occupancy. It may therefore +overtake ordinary messages that have not yet reached the transport. Tests or +applications that need to cancel work already running on a peer must establish +that the peer has started it, rather than relying on a synchronous `send_request` +call or a scheduler yield. + ## Interoperability Protocol-level (`$/`-prefixed) notifications are optional by design. The SDK diff --git a/md/transport-architecture.md b/md/transport-architecture.md index ef9b84fe..9f7764dc 100644 --- a/md/transport-architecture.md +++ b/md/transport-architecture.md @@ -29,7 +29,7 @@ The architecture separates two distinct responsibilities: This separation enables: -- **In-process efficiency**: Components in the same process can skip serialization +- **In-process efficiency**: Components can pass frames without a serialize/parse round trip - **Transport flexibility**: Easy to add new transport types (WebSockets, named pipes, etc.) - **Testability**: Mock transports for unit testing - **Clarity**: Clear boundaries between protocol and I/O concerns @@ -55,8 +55,8 @@ At that boundary: - **Below**: Transport actors parse and serialize JSON-RPC frames - **Boundary**: `TransportFrame` carries one raw message, a structurally non-empty batch, or a malformed wire value retained for a relay -- **In-process API**: `Channel::rx` and `Channel::tx` carry `TransportFrame` - directly, so adapters cannot accidentally flatten a batch +- **In-process API**: `Channel::rx` and `Channel::tx` carry `BudgetedFrame` + envelopes containing the complete `TransportFrame` and its byte permit - **Failures**: I/O and connection failures are returned by the future driving a transport; they are not sent as channel entries @@ -76,8 +76,8 @@ These actors live in the protocol connection core and understand JSON-RPC semant #### Outgoing Protocol Actor ``` -Input: mpsc::UnboundedReceiver -Output: mpsc::UnboundedSender +Input: Bounded application admission queues +Output: FrameSender (BudgetedFrame) ``` Responsibilities: @@ -89,7 +89,7 @@ Responsibilities: #### Incoming Protocol Actor ``` -Input: mpsc::UnboundedReceiver +Input: FrameReceiver (BudgetedFrame) Output: Routes to pending request awaiters or registered handlers ``` @@ -115,6 +115,18 @@ The shared pending-reply registry manages request/response correlation: Runs user-spawned concurrent tasks via `cx.spawn()`. +Admission counts queued and running tasks together, using +`ConnectionLimits::max_queued_frames`. Accepted tasks are polled concurrently; +there is no second waiting pool behind permanent child connection drivers. +When all live slots are occupied, spawning or registering an ordered response +consumer fails immediately instead of accepting work that cannot make progress. +The connection's own transport driver runs outside this task pool. + +Native MCP supervisors also register protected cleanup acknowledgments. +Connection shutdown signals their cancellation and continues driving them and +their scoped tool runners until cleanup finishes, including when another task +or transport fails. Unrelated user tasks are not joined indefinitely. + ### Transport Actors These actors are driven by physical transport components. They understand @@ -124,7 +136,7 @@ correlate responses with pending requests: #### Transport Outgoing Actor ``` -Input: mpsc::UnboundedReceiver +Input: FrameReceiver (BudgetedFrame) Output: Writes to I/O (byte stream, channel, socket, etc.) ``` @@ -135,13 +147,13 @@ For byte streams: For in-process channels: -- Directly forward `TransportFrame` to the channel +- Forward `BudgetedFrame` to preserve both the frame and its admission #### Transport Incoming Actor ``` Input: Reads from I/O (byte stream, channel, socket, etc.) -Output: mpsc::UnboundedSender +Output: FrameSender (BudgetedFrame) ``` For byte streams: @@ -156,7 +168,7 @@ For byte streams: For in-process channels: -- Directly forward `TransportFrame` from the channel +- Forward `BudgetedFrame` from the channel without releasing admission The public `Channel` boundary preserves complete frames. The SDK continues to initiate requests and notifications as individual JSON-RPC messages; response @@ -219,7 +231,7 @@ Outgoing Protocol Actor | - Subscribe to replies | - Convert to RawJsonRpcMessage v - | TransportFrame (single message or batch response) + | BudgetedFrame (single message or batch response, with admission) | Transport Outgoing Actor | - Serialize (byte streams) @@ -237,7 +249,7 @@ Transport Incoming Actor | - Parse (byte streams) | - Or forward directly (channels) v - | TransportFrame (single message or incoming batch) + | BudgetedFrame (single message or incoming batch, with admission) | Incoming Protocol Actor | - Route responses → pending request awaiters @@ -286,6 +298,27 @@ an intermediate copy. A forwarded frame keeps its permit through any adapter queue, deferred dispatch, or writer. This accounting is internal and does not change the JSON-RPC wire shape. +### Bounded admission + +`Channel::duplex_with_limits` accepts `ConnectionLimits`. Defaults are a 16 MiB +maximum frame, a 64 MiB shared duplex serialized-payload budget, and 32 queued +frames per direction. Responses and cancellation have reserved byte capacity. +These are serialized-data and item bounds, not an exact bound on allocator +overhead or application-owned memory. + +Frame-sink clones share item capacity, including slots reserved by +`Sink::poll_ready`. `try_send` fails immediately at capacity; async sends wait +outside the dispatcher. Receiver dequeue releases the queue slot, but the +`BudgetedFrame` keeps its byte charge through deferred processing and writing. +Forward that envelope intact. Forwarding between independently budgeted +channels must also satisfy the destination's limits. + +When retaining only metadata derived from a frame, use +`FramePermit::try_reserve_metadata` with its measured serialized size. This +reserves an independent charge in every source budget; it fails immediately +rather than waiting on the payload's own reservation. Drop the original permit +once the payload is consumed, and retain the new one with the metadata. + ## Transport Implementations ### Byte Stream Transport @@ -317,13 +350,14 @@ Use cases: ### In-Process Channel For components in the same process, `Channel::duplex()` creates paired -endpoints and skips serialization entirely. Relays forward each received -`TransportFrame` without unpacking it; this preserves batch boundaries and the -original representation of malformed wire input. +endpoints without encoding and reparsing a wire message between components. +Relays forward each received `BudgetedFrame` without unpacking it; this preserves +batch boundaries, admission, and the representation of malformed wire input. +Admission still measures serialized size to enforce the shared byte budget. Benefits: -- **Zero serialization overhead**: Messages passed by value +- **No wire round trip**: Frames are passed by value, with serialized-size accounting - **Same-process efficiency**: Ideal for conductor with in-process proxies - **Explicit wire state**: No serialize/parse round trip is required, while a malformed value received from a physical transport remains an explicit frame diff --git a/src/agent-client-protocol-conductor/tests/initialization_v2.rs b/src/agent-client-protocol-conductor/tests/initialization_v2.rs index cd9a0fdc..2741113f 100644 --- a/src/agent-client-protocol-conductor/tests/initialization_v2.rs +++ b/src/agent-client-protocol-conductor/tests/initialization_v2.rs @@ -1007,7 +1007,7 @@ async fn v2_proxy_session_helper_reissues_cancellation_for_the_downstream_hop() let (editor_out, conductor_in) = duplex(4096); let (conductor_out, editor_in) = duplex(4096); let transport = ByteStreams::new(editor_out.compat_write(), editor_in.compat()); - let client_request_id = tokio::time::timeout( + let (client_request_id, parked_id) = tokio::time::timeout( std::time::Duration::from_secs(10), Client .v2() @@ -1027,6 +1027,10 @@ async fn v2_proxy_session_helper_reissues_cancellation_for_the_downstream_hop() let pending = cx.send_request(v2::NewSessionRequest::new("/park-session")); let client_request_id = pending.id().clone(); + let parked_id = parked_id_rx + .next() + .await + .ok_or_else(|| Error::internal_error().data("parked request channel closed"))?; pending.cancel()?; let error = pending .block_task() @@ -1039,16 +1043,12 @@ async fn v2_proxy_session_helper_reissues_cancellation_for_the_downstream_hop() .block_task() .await?; assert_eq!(response.session_id, v2::SessionId::new("normal-session")); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }), ) .await .expect("v2 proxy cancellation test timed out")?; - let parked_id = tokio::time::timeout(std::time::Duration::from_secs(2), parked_id_rx.next()) - .await - .expect("agent should observe the forwarded request") - .ok_or_else(|| Error::internal_error().data("parked request channel closed"))?; assert_ne!( parked_id, client_request_id, "each proxy hop must allocate its own request ID" @@ -1120,7 +1120,7 @@ async fn v2_proxy_resume_helper_reissues_cancellation_for_the_downstream_hop() - let (editor_out, conductor_in) = duplex(4096); let (conductor_out, editor_in) = duplex(4096); let transport = ByteStreams::new(editor_out.compat_write(), editor_in.compat()); - let client_request_id = tokio::time::timeout( + let (client_request_id, parked_id) = tokio::time::timeout( std::time::Duration::from_secs(10), Client .v2() @@ -1143,6 +1143,10 @@ async fn v2_proxy_resume_helper_reissues_cancellation_for_the_downstream_hop() - "/park-session", )); let client_request_id = pending.id().clone(); + let parked_id = parked_id_rx + .next() + .await + .ok_or_else(|| Error::internal_error().data("parked request channel closed"))?; pending.cancel()?; let error = pending .block_task() @@ -1158,16 +1162,12 @@ async fn v2_proxy_resume_helper_reissues_cancellation_for_the_downstream_hop() - .block_task() .await?; assert_eq!(response, v2::ResumeSessionResponse::new()); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }), ) .await .expect("v2 resume proxy cancellation test timed out")?; - let parked_id = tokio::time::timeout(std::time::Duration::from_secs(2), parked_id_rx.next()) - .await - .expect("agent should observe the forwarded resume request") - .ok_or_else(|| Error::internal_error().data("parked request channel closed"))?; assert_ne!( parked_id, client_request_id, "each proxy hop must allocate its own request ID" diff --git a/src/agent-client-protocol-conductor/tests/request_cancellation.rs b/src/agent-client-protocol-conductor/tests/request_cancellation.rs index 19ca45fc..a2449cfd 100644 --- a/src/agent-client-protocol-conductor/tests/request_cancellation.rs +++ b/src/agent-client-protocol-conductor/tests/request_cancellation.rs @@ -204,12 +204,12 @@ async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Erro .await }); - let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let initialize = cx .send_request(InitializeRequest::new(ProtocolVersion::V1)) .block_task() @@ -220,6 +220,7 @@ async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Erro message: "park".into(), }); let client_request_id = request.id().clone(); + let parked_id = next_with_timeout(&mut parked_id_rx).await; request.cancel()?; // The cancellation reaches the agent hop by hop, and the @@ -242,7 +243,7 @@ async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Erro .await?; assert_eq!(barrier.result, "echo: barrier"); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }, ) .await @@ -250,10 +251,10 @@ async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Erro .await .expect("test timed out") .expect("client failed"); + let (client_request_id, parked_id) = client_result; // The agent saw exactly one `$/cancel_request`, for the request ID on // its own connection. - 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" @@ -277,6 +278,9 @@ async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Erro let (client_cancel_tx, mut client_cancel_rx) = mpsc::unbounded(); // The JSON-RPC id of the parked request, as seen by the client. let (parked_id_tx, mut parked_id_rx) = mpsc::unbounded(); + let (parked_tx, parked_rx) = tokio::sync::oneshot::channel(); + let parked_tx = Arc::new(Mutex::new(Some(parked_tx))); + let parked_rx = Arc::new(Mutex::new(Some(parked_rx))); let agent = Agent .builder() @@ -287,11 +291,12 @@ async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Erro agent_client_protocol::on_receive_request!(), ) .on_receive_request( - async |request: SimpleRequest, - responder: Responder, - cx: ConnectionTo| { + async move |request: SimpleRequest, + responder: Responder, + cx: ConnectionTo| { if request.message == "trigger reverse cancel" { let connection = cx.clone(); + let parked_rx = parked_rx.lock().unwrap().take().expect("one trigger"); cx.spawn(async move { // Send a request to the client, cancel it, and report // how it concluded as the response to the trigger. @@ -299,6 +304,10 @@ async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Erro connection.send_request(SimpleRequest { message: "park".into(), }); + tokio::time::timeout(Duration::from_secs(10), parked_rx) + .await + .expect("timed out waiting for client to park request") + .expect("client closed parked request channel"); upstream.cancel()?; let error = upstream .block_task() @@ -342,6 +351,13 @@ async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Erro cx: ConnectionTo| { assert_eq!(request.message, "park"); parked_id_tx.unbounded_send(responder.id().clone()).unwrap(); + parked_tx + .lock() + .unwrap() + .take() + .expect("one parked request") + .send(()) + .expect("agent still waiting to cancel"); let cancellation = responder.cancellation(); cx.spawn(async move { let response = cancellation @@ -534,7 +550,7 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E .await }); - let client_prompt_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .on_receive_request( @@ -579,7 +595,7 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E ) .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let initialize = cx .send_request(InitializeRequest::new(ProtocolVersion::V1)) .block_task() @@ -598,6 +614,8 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E vec!["park".into()], )); let client_prompt_id = prompt.id().clone(); + let prompt_id = next_with_timeout(&mut prompt_id_rx).await; + let permission_id = next_with_timeout(&mut permission_id_rx).await; prompt.cancel()?; let error = prompt @@ -619,7 +637,7 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E .await?; assert_eq!(barrier.stop_reason, StopReason::EndTurn); - Ok(client_prompt_id) + Ok((client_prompt_id, prompt_id, permission_id)) }, ) .await @@ -627,10 +645,10 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E .await .expect("test timed out") .expect("client failed"); + let (client_prompt_id, prompt_id, permission_id) = client_result; // The agent saw exactly one `$/cancel_request` (for the prompt), with the // ID of the prompt on the conductor-to-agent connection. - 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" @@ -641,7 +659,6 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E // The client saw exactly one `$/cancel_request` (for the permission // request), with the ID of that request on the client's own connection. - 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); @@ -717,12 +734,12 @@ async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error .await }); - let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let initialize = cx .send_request(InitializeRequest::new(ProtocolVersion::V1)) .block_task() @@ -732,6 +749,7 @@ async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error let request: SentRequest = 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 @@ -750,7 +768,7 @@ async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error .await?; assert_eq!(session.session_id, SessionId::new("normal-session")); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }, ) .await @@ -758,10 +776,10 @@ async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error .await .expect("test timed out") .expect("client failed"); + let (client_request_id, parked_id) = client_result; // The agent saw exactly one `$/cancel_request`, for the `session/new` ID // on its own connection. - 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" @@ -845,12 +863,12 @@ async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), E .await }); - let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let initialize = cx .send_request(InitializeRequest::new(ProtocolVersion::V1)) .block_task() @@ -860,6 +878,7 @@ async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), E let request: SentRequest = 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 @@ -876,7 +895,7 @@ async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), E .await?; assert_eq!(session.session_id, SessionId::new("normal-session")); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }, ) .await @@ -884,8 +903,8 @@ async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), E .await .expect("test timed out") .expect("client failed"); + let (client_request_id, parked_id) = client_result; - 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" @@ -1064,6 +1083,7 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - let request: SentRequest = 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 @@ -1081,7 +1101,7 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - assert_eq!(session.session_id, SessionId::new("normal-session")); let probe_barrier = next_with_timeout(&mut probe_barrier_rx).await; - Ok((client_request_id, probe_barrier)) + Ok((client_request_id, parked_id, probe_barrier)) }, ) .await @@ -1089,9 +1109,8 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - .await .expect("test timed out") .expect("client failed"); - let (client_request_id, probe_barrier) = client_result; + let (client_request_id, parked_id, probe_barrier) = client_result; - 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" @@ -1404,15 +1423,16 @@ async fn initialize_cancellation_propagates_through_proxy() -> Result<(), Error> .await }); - let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let request: SentRequest = cx.send_request(InitializeRequest::new(ProtocolVersion::V1)); let client_request_id = request.id().clone(); + let parked_id = next_with_timeout(&mut parked_id_rx).await; request.cancel()?; let error = request @@ -1429,7 +1449,7 @@ async fn initialize_cancellation_propagates_through_proxy() -> Result<(), Error> .await?; assert_eq!(initialize.protocol_version, ProtocolVersion::V1); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }, ) .await @@ -1437,10 +1457,10 @@ async fn initialize_cancellation_propagates_through_proxy() -> Result<(), Error> .await .expect("test timed out") .expect("client failed"); + let (client_request_id, parked_id) = client_result; // The agent saw exactly one `$/cancel_request`, for the `initialize` ID // on its own connection. - 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" diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index 14b92f11..791172f0 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -215,6 +215,7 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { .start_sse( Some(session_id), sse_event_tx.clone(), + false, SseStartContext { events: &mut sse_event_rx, outgoing: &mut outgoing, @@ -265,11 +266,25 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { // must not be blocked behind the request they answer. Ok((post, session_ids)) => { state.attach_pending_permits(&post.pending_requests, &permit); + if let Err(error) = + check_post_capacity(&posts, max_operations, bypass_ordered) + { + break 'transport Err(error); + } + if bypass_ordered { + posts.responses.push_budgeted(post, permit); + } else { + posts.ordered.push_budgeted(post, permit); + } + // The POST registers bounded session mailboxes. Poll it + // while establishing the GET, rather than waiting for + // a GET that cannot succeed before the POST. for session_id in session_ids { match lifecycle .start_sse( Some(session_id), sse_event_tx.clone(), + true, SseStartContext { events: &mut sse_event_rx, outgoing: &mut outgoing, @@ -281,22 +296,17 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { .await { Ok(SseStartOutcome::Established) => {} + Ok(SseStartOutcome::OutgoingClosed) + if buffered_outgoing.is_empty() && posts.is_empty() => + { + break 'transport Ok(()); + } Ok(SseStartOutcome::OutgoingClosed) => { break 'transport Err(sse_setup_blocked_output_error()); } Err(error) => break 'transport Err(error), } } - if let Err(error) = - check_post_capacity(&posts, max_operations, bypass_ordered) - { - break 'transport Err(error); - } - if bypass_ordered { - posts.responses.push_budgeted(post, permit); - } else { - posts.ordered.push_budgeted(post, permit); - } } Err(error) => { error!("POST failed: {error}"); @@ -318,6 +328,7 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { .start_sse( None, sse_event_tx.clone(), + false, SseStartContext { events: &mut sse_event_rx, outgoing: &mut outgoing, @@ -347,12 +358,34 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { continue; } - if let Some(session_id) = session_id_from_message(&msg) { + let session_id = session_id_from_message(&msg); + if let Err(error) = check_post_capacity(&posts, max_operations, bypass_ordered) { + break Err(error); + } + match state.prepare_post(msg) { + // Responses and cancellation must not be blocked behind a POST + // that may itself be waiting for their delivery. + Ok(post) => { + state.attach_pending_permits(&post.pending_requests, &permit); + if bypass_ordered { + posts.responses.push_budgeted(post, permit); + } else { + posts.ordered.push_budgeted(post, permit); + } + } + Err(e) => { + error!("POST failed: {e}"); + break Err(AcpError::internal_error().data(format!("POST: {e}"))); + } + } + + if let Some(session_id) = session_id { for session_id in state.register_session_streams([session_id]) { match lifecycle .start_sse( Some(session_id), sse_event_tx.clone(), + true, SseStartContext { events: &mut sse_event_rx, outgoing: &mut outgoing, @@ -364,6 +397,11 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { .await { Ok(SseStartOutcome::Established) => {} + Ok(SseStartOutcome::OutgoingClosed) + if buffered_outgoing.is_empty() && posts.is_empty() => + { + break 'transport Ok(()); + } Ok(SseStartOutcome::OutgoingClosed) => { break 'transport Err(sse_setup_blocked_output_error()); } @@ -371,26 +409,6 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { } } } - - if let Err(error) = check_post_capacity(&posts, max_operations, bypass_ordered) { - break Err(error); - } - match state.prepare_post(msg) { - // Responses and cancellation must not be blocked behind a POST - // that may itself be waiting for their delivery. - Ok(post) => { - state.attach_pending_permits(&post.pending_requests, &permit); - if bypass_ordered { - posts.responses.push_budgeted(post, permit); - } else { - posts.ordered.push_budgeted(post, permit); - } - } - Err(e) => { - error!("POST failed: {e}"); - break Err(AcpError::internal_error().data(format!("POST: {e}"))); - } - } }; lifecycle.close().await; @@ -645,6 +663,7 @@ impl HttpTransportLifecycle { &mut self, session_id: Option, event_tx: mpsc::Sender, + wait_for_post: bool, context: SseStartContext<'_>, ) -> Result { let SseStartContext { @@ -655,15 +674,32 @@ impl HttpTransportLifecycle { state, } = context; let mut establishing = FuturesUnordered::new(); - establishing.push(self.begin_sse(session_id, event_tx.clone())?); + let mut session_id = session_id; + if !wait_for_post { + establishing.push(self.begin_sse(session_id.take(), event_tx.clone())?); + } loop { - if establishing.is_empty() { + // A session-scoped POST receives 202 only after its route and + // mailbox are registered. Keep pumping all existing SSE streams, + // callbacks, and POST completions until that admission completes; + // then open the session GET without racing a 409. + if wait_for_post && posts.is_empty() && session_id.is_some() { + establishing.push(self.begin_sse(session_id.take(), event_tx.clone())?); + } + if establishing.is_empty() && session_id.is_none() { return Ok(SseStartOutcome::Established); } let outcome = { let failure = self.sse_tasks.next_failure().fuse(); - let established_next = establishing.next().fuse(); + let established_next = async { + if establishing.is_empty() { + futures::future::pending().await + } else { + establishing.next().await + } + } + .fuse(); let sse_event_next = events.next().fuse(); let outgoing_next = outgoing.next().fuse(); let ordered_post_next = posts.ordered.next_completion().fuse(); @@ -691,7 +727,9 @@ impl HttpTransportLifecycle { return Err(sse_failure_error(self.sse_tasks.next_failure().await)); } SseStartWait::Established(None) => { - return Ok(SseStartOutcome::Established); + if session_id.is_none() { + return Ok(SseStartOutcome::Established); + } } SseStartWait::Failure(failure) => return Err(sse_failure_error(failure)), SseStartWait::SseEvent(Some(event)) => { @@ -2133,7 +2171,6 @@ mod tests { let post_count = Arc::new(AtomicUsize::new(0)); let emit_response = Arc::new(Notify::new()); let connection_stream_established = Arc::new(AtomicBool::new(false)); - let source_stream_established = Arc::new(AtomicBool::new(false)); let response_batch = json!([ { "jsonrpc": "2.0", @@ -2146,20 +2183,16 @@ mod tests { post({ let post_count = post_count.clone(); let connection_stream_established = connection_stream_established.clone(); - let source_stream_established = source_stream_established.clone(); move |body: String| { let post_count = post_count.clone(); let post_tx = post_tx.clone(); let connection_stream_established = connection_stream_established.clone(); - let source_stream_established = source_stream_established.clone(); async move { if post_count.fetch_add(1, Ordering::SeqCst) == 0 { return initialize_response().await.into_response(); } - if !connection_stream_established.load(Ordering::SeqCst) - || !source_stream_established.load(Ordering::SeqCst) - { + if !connection_stream_established.load(Ordering::SeqCst) { return StatusCode::CONFLICT.into_response(); } post_tx @@ -2173,13 +2206,13 @@ mod tests { let emit_response = emit_response.clone(); let response_batch = response_batch.clone(); let connection_stream_established = connection_stream_established.clone(); - let source_stream_established = source_stream_established.clone(); + let post_count = post_count.clone(); move |headers: HeaderMap| { let emit_response = emit_response.clone(); let response_batch = response_batch.clone(); let get_tx = get_tx.clone(); let connection_stream_established = connection_stream_established.clone(); - let source_stream_established = source_stream_established.clone(); + let post_count = post_count.clone(); async move { let session_id = headers .get(HEADER_SESSION_ID) @@ -2187,13 +2220,15 @@ mod tests { .map(String::from); let is_connection_stream = session_id.is_none(); let is_source_stream = session_id.as_deref() == Some("source-session"); + if is_source_stream && post_count.load(Ordering::SeqCst) < 2 { + return StatusCode::CONFLICT.into_response(); + } if is_connection_stream { sleep(Duration::from_millis(50)).await; connection_stream_established.store(true, Ordering::SeqCst); } if is_source_stream { sleep(Duration::from_millis(50)).await; - source_stream_established.store(true, Ordering::SeqCst); } get_tx.send(session_id).unwrap(); @@ -2206,7 +2241,7 @@ mod tests { } futures::future::pending::<()>().await; }; - Sse::new(stream) + Sse::new(stream).into_response() } } }) @@ -2879,6 +2914,7 @@ mod tests { lifecycle.start_sse( Some("later-session".to_string()), event_tx, + true, SseStartContext { events: &mut event_rx, outgoing: &mut outgoing, @@ -2901,6 +2937,15 @@ mod tests { #[tokio::test] async fn stalled_sse_establishment_keeps_callback_responses_moving() { + check_callback_response_progress(false).await; + } + + #[tokio::test] + async fn cold_session_post_admission_keeps_callback_responses_moving() { + check_callback_response_progress(true).await; + } + + async fn check_callback_response_progress(wait_for_post: bool) { let release_get = Arc::new(Notify::new()); let complete_earlier_post = Arc::new(Notify::new()); let app = Router::new().route( @@ -3004,6 +3049,7 @@ mod tests { lifecycle.start_sse( Some("later-session".to_string()), event_tx, + wait_for_post, SseStartContext { events: &mut event_rx, outgoing: &mut outgoing, diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index 23b5591b..f38504c9 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -92,6 +92,9 @@ impl OutboundMailbox { impl OutboundLease { pub(crate) async fn recv(&mut self) -> Option { + // The previously returned text has been handed to the transport. Do + // not retain its byte charge while waiting for the next frame. + self.current.take(); let value = self .receiver .as_mut() @@ -103,6 +106,7 @@ impl OutboundLease { } pub(crate) fn try_recv(&mut self) -> Result { + self.current.take(); let value = self .receiver .as_mut() @@ -143,6 +147,10 @@ impl Connection { self.send_budgeted_frame_to_agent(frame) } + pub(crate) fn agent_channel_closed(&self) -> bool { + self.inbound_tx.is_closed() + } + pub(crate) fn admit_frame_to_agent( &self, frame: TransportFrame, @@ -412,10 +420,21 @@ impl HttpOutbound { if new_sessions.len().saturating_add(routes.len()) > available { return Err("HTTP pending route or session capacity exceeded"); } - for id in &new_sessions { + // Reserve every new ID before publishing any metadata. Session + // mailboxes retain only their key, not the unrelated POST payload; + // pending response routes still retain their originating frame. + let session_permits = new_sessions + .iter() + .map(|id| { + permit + .try_reserve_metadata(session_metadata_bytes(id)) + .map_err(|_| "HTTP session metadata capacity exceeded") + }) + .collect::, _>>()?; + for (id, session_permit) in new_sessions.iter().zip(session_permits) { streams.insert( id.clone(), - (Arc::new(OutboundMailbox::new()), Some(permit.clone())), + (Arc::new(OutboundMailbox::new()), Some(session_permit)), ); } for (id, route) in routes { @@ -488,8 +507,14 @@ impl HttpOutbound { if streams.len().saturating_add(pending_count) >= self.limits.max_queued_frames.max(1) { return Err("HTTP session stream capacity exceeded"); } + let session_permit = permit + .try_reserve_metadata(session_metadata_bytes(session_id)) + .map_err(|_| "HTTP session metadata capacity exceeded")?; let stream = Arc::new(OutboundMailbox::new()); - streams.insert(session_id.to_string(), (stream.clone(), Some(permit))); + streams.insert( + session_id.to_string(), + (stream.clone(), Some(session_permit)), + ); Ok(stream) } @@ -807,6 +832,14 @@ fn pending_route_key(id: &RequestId) -> Option { } } +fn session_metadata_bytes(session_id: &str) -> usize { + // The retained key is one JSON string; account for escaping and quotes, + // not for the unrelated source request/response payload. + serde_json::to_string(session_id) + .expect("string serialization cannot fail") + .len() +} + fn response_session_id(msg: &RawJsonRpcMessage) -> Option<&str> { let RawJsonRpcMessage::Response(RpcResponse::Result { result, .. }) = msg else { return None; diff --git a/src/agent-client-protocol-http/src/connection_admission_tests.rs b/src/agent-client-protocol-http/src/connection_admission_tests.rs index 23066f94..dfe70f65 100644 --- a/src/agent-client-protocol-http/src/connection_admission_tests.rs +++ b/src/agent-client-protocol-http/src/connection_admission_tests.rs @@ -3,6 +3,105 @@ use serde_json::json; use super::*; +#[tokio::test] +async fn outbound_lease_releases_delivered_frame_before_waiting_or_idle_poll() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".into(), json!({"data": "x".repeat(100)})) + .unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes + 1, + max_queued_bytes: bytes * 2 + 1, + max_queued_frames: 2, + }; + let (_, channel) = Channel::duplex_with_limits(limits); + let admission = channel.tx.admission(); + let mailbox = OutboundMailbox::new(); + let mut lease = mailbox.try_acquire().unwrap(); + let (_, first) = admission.try_admit(frame.clone()).unwrap().into_parts(); + mailbox + .push_with_permit("first".into(), Some(first)) + .unwrap(); + assert_eq!(lease.recv().await.as_deref(), Some("first")); + assert!(admission.try_admit(frame.clone()).is_err()); + assert!(lease.try_recv().is_err()); + let (_, second) = admission.try_admit(frame.clone()).unwrap().into_parts(); + mailbox + .push_with_permit("second".into(), Some(second)) + .unwrap(); + assert_eq!(lease.try_recv().unwrap(), "second"); + assert!(admission.try_admit(frame.clone()).is_err()); + let wait = tokio::spawn(async move { lease.recv().await }); + tokio::task::yield_now().await; + let (_, third) = admission.try_admit(frame).unwrap().into_parts(); + mailbox + .push_with_permit("third".into(), Some(third)) + .unwrap(); + assert_eq!(wait.await.unwrap().as_deref(), Some("third")); +} + +#[tokio::test] +async fn session_key_does_not_pin_unrelated_post_payload() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification( + "session/update".into(), + json!({"sessionId": "persisted", "payload": "x".repeat(512)}), + ) + .unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2 + 64, + max_queued_frames: 2, + }; + let (_, channel) = Channel::duplex_with_limits(limits); + let admission = channel.tx.admission(); + let (_, source) = admission.try_admit(frame.clone()).unwrap().into_parts(); + let mut http = HttpOutbound::new(); + http.limits = limits; + http.register_post_routes(&["persisted".into()], &[], &source) + .await + .unwrap(); + drop(source); + assert!(http.session_streams.read().await.contains_key("persisted")); + assert!( + admission.try_admit(frame).is_ok(), + "the retained session key must not pin its source payload" + ); +} + +#[tokio::test] +async fn failed_session_key_reservation_does_not_publish_partial_batch() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".into(), json!({"payload": "x".repeat(512)})) + .unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2 + 12, + max_queued_frames: 3, + }; + let (_, channel) = Channel::duplex_with_limits(limits); + let (_, source) = channel + .tx + .admission() + .try_admit(frame) + .unwrap() + .into_parts(); + let mut http = HttpOutbound::new(); + http.limits = limits; + assert!( + http.register_post_routes(&["one".into(), "another".into()], &[], &source) + .await + .is_err() + ); + assert!(http.session_streams.read().await.is_empty()); + assert!(http.pending_routes.lock().await.is_empty()); +} + #[tokio::test] async fn route_and_session_metadata_admission_is_atomic_and_releases_permits() { let frame = TransportFrame::Single( @@ -10,8 +109,8 @@ async fn route_and_session_metadata_admission_is_atomic_and_releases_permits() { ); let bytes = frame.to_json().unwrap().len(); let limits = ConnectionLimits { - max_frame_bytes: bytes + 128, - max_queued_bytes: bytes * 2 + 128, + max_frame_bytes: bytes + 16, + max_queued_bytes: bytes * 2 + 32, max_queued_frames: 2, }; let (_caller, transport) = Channel::duplex_with_limits(limits); @@ -46,7 +145,10 @@ async fn route_and_session_metadata_admission_is_atomic_and_releases_permits() { ), Some(ResponseRoute::Session("one".into())) ); - assert!(admission.try_admit(frame.clone()).is_err()); + assert!( + admission.try_admit(frame.clone()).is_ok(), + "removing a pending route releases its full source-frame charge" + ); http.session_streams.write().await.clear(); assert!(admission.try_admit(frame).is_ok()); } diff --git a/src/agent-client-protocol-http/src/server.rs b/src/agent-client-protocol-http/src/server.rs index 695b1699..07cab781 100644 --- a/src/agent-client-protocol-http/src/server.rs +++ b/src/agent-client-protocol-http/src/server.rs @@ -196,9 +196,190 @@ async fn handle_get( #[cfg(test)] mod tests { use super::*; + use agent_client_protocol::{ + Channel, ConnectTo, RawJsonRpcMessage, TransportBatch, TransportFrame, + schema::v1::RequestId, + }; use axum::body::Body; + use futures::{StreamExt, future::BoxFuture}; + use serde_json::json; + use tokio::{ + net::TcpListener, + time::{Duration, timeout}, + }; use tower::{Layer as _, ServiceExt as _, service_fn}; + struct HistoryAgent; + + impl crate::connection::AgentFactory for HistoryAgent { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (mut agent, transport) = Channel::duplex(); + let run = Box::pin(async move { + while let Some(frame) = agent.rx.next().await { + let messages = match frame.into_frame() { + TransportFrame::Single(message) => vec![message], + TransportFrame::Batch(batch) => batch + .entries() + .filter_map(|entry| match entry { + agent_client_protocol::TransportBatchEntry::Message(message) => { + Some(message.clone()) + } + agent_client_protocol::TransportBatchEntry::Malformed { + .. + } => None, + }) + .collect(), + TransportFrame::Malformed { .. } => continue, + }; + for message in messages { + let RawJsonRpcMessage::Request(request) = message else { + continue; + }; + if request.method.as_ref() != "initialize" { + let Some(agent_client_protocol::RawJsonRpcParams::Object(params)) = + request.params.as_ref() + else { + panic!("session request must have object params"); + }; + for index in 0..2 { + agent + .tx + .send_frame(TransportFrame::Single( + RawJsonRpcMessage::notification( + "session/update".into(), + json!({"sessionId": params["sessionId"], "index": index}), + ) + .unwrap(), + )) + .await + .unwrap(); + } + } + agent + .tx + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( + request.id, + Ok(json!({})), + ))) + .await + .unwrap(); + } + } + Ok(()) + }); + (transport, run) + } + } + + #[tokio::test] + async fn cold_session_post_registers_stream_before_history_for_single_and_batch() { + let registry = Arc::new(ConnectionRegistry::new(Arc::new(HistoryAgent))); + let app = AcpHttpServer { + registry: registry.clone(), + options: ServerOptions::default(), + } + .into_router(); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + for (index, method) in ["session/load", "session/resume"].into_iter().enumerate() { + let client = crate::client::HttpClient::new(format!("http://{address}")).unwrap(); + let (mut caller, driver) = client.into_channel_and_future(); + let driver = tokio::spawn(driver); + caller + .tx + .send_frame(TransportFrame::Single( + RawJsonRpcMessage::request( + "initialize".into(), + json!({}), + RequestId::Number(1), + ) + .unwrap(), + )) + .await + .unwrap(); + let init = timeout(Duration::from_secs(2), caller.rx.next()) + .await + .unwrap() + .unwrap(); + assert!(matches!( + init.frame(), + TransportFrame::Single(RawJsonRpcMessage::Response(_)) + )); + drop(init); + let request = RawJsonRpcMessage::request( + method.into(), + json!({"sessionId": "persisted"}), + RequestId::Number(2), + ) + .unwrap(); + let frame = if index == 0 { + TransportFrame::Single(request) + } else { + let second = RawJsonRpcMessage::request( + method.into(), + json!({"sessionId": "other-persisted"}), + RequestId::Number(3), + ) + .unwrap(); + TransportFrame::Batch(TransportBatch::from_messages([request, second]).unwrap()) + }; + caller.tx.send_frame(frame).await.unwrap(); + let mut seen = std::collections::BTreeMap::>::new(); + let session_count = index + 1; + let mut responses = 0; + for _ in 0..session_count * 3 { + let frame = timeout(Duration::from_secs(3), caller.rx.next()) + .await + .unwrap() + .unwrap(); + match frame.frame() { + TransportFrame::Single(RawJsonRpcMessage::Notification(notification)) => { + let agent_client_protocol::RawJsonRpcParams::Object(params) = + notification.params.as_ref().unwrap() + else { + panic!("history update must have object params"); + }; + seen.entry(params["sessionId"].as_str().unwrap().to_owned()) + .or_default() + .push(params["index"].as_u64().unwrap()); + } + TransportFrame::Single(RawJsonRpcMessage::Response(_)) => { + let response: serde_json::Value = + serde_json::from_str(&frame.frame().to_json().unwrap()).unwrap(); + let session = match response["id"].as_u64().unwrap() { + 2 => "persisted", + 3 => "other-persisted", + id => panic!("unexpected response ID: {id}"), + }; + assert_eq!( + seen.get(session).map(Vec::as_slice), + Some([0, 1].as_slice()), + "{method} must deliver history before its response" + ); + responses += 1; + } + other => panic!("unexpected history frame: {other:?}"), + } + } + assert_eq!(seen.len(), session_count); + assert_eq!(responses, session_count); + drop(caller); + timeout(Duration::from_secs(3), driver) + .await + .unwrap() + .unwrap() + .unwrap(); + } + assert_eq!(registry.len().await, 0); + server.abort(); + } + #[test] fn cors_is_disabled_by_default() { assert_eq!(ServerOptions::default().cors, CorsOptions::Disabled); diff --git a/src/agent-client-protocol-http/src/websocket_server.rs b/src/agent-client-protocol-http/src/websocket_server.rs index 83882b25..4e614aab 100644 --- a/src/agent-client-protocol-http/src/websocket_server.rs +++ b/src/agent-client-protocol-http/src/websocket_server.rs @@ -162,13 +162,23 @@ where { trace!(connection_id = %connection_id, session_id = %sid, request_id = ?req.id, "Client → Agent (session)"); } - if connection.send_frame_to_agent(frame).is_err() { - error!(connection_id = %connection_id, "Agent channel closed"); - drain_outbound_until_closed(ws_tx, outbound_rx, closed, connection_id).await; - false - } else { - true + let frame = match connection.admit_frame_to_agent(frame) { + Ok(frame) => frame, + Err(error) => { + warn!(connection_id = %connection_id, "Rejecting WebSocket frame: {error}"); + return false; + } + }; + if let Err(error) = connection.send_budgeted_frame_to_agent(frame) { + if connection.agent_channel_closed() { + error!(connection_id = %connection_id, "Agent channel closed"); + drain_outbound_until_closed(ws_tx, outbound_rx, closed, connection_id).await; + } else { + warn!(connection_id = %connection_id, "Rejecting WebSocket frame: {error}"); + } + return false; } + true } async fn drain_outbound_until_closed( @@ -296,6 +306,161 @@ mod tests { } } + struct LimitedAgentFactory { + forwarded: mpsc::UnboundedSender, + } + + struct StalledAgentFactory; + + impl AgentFactory for StalledAgentFactory { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (agent, transport) = + Channel::duplex_with_limits(agent_client_protocol::ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 2048, + max_queued_frames: 4, + }); + let future = Box::pin(async move { + std::future::pending::<()>().await; + drop(agent); + Ok(()) + }); + (transport, future) + } + } + + impl AgentFactory for LimitedAgentFactory { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let limits = agent_client_protocol::ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 2048, + max_queued_frames: 4, + }; + let (mut agent, transport) = Channel::duplex_with_limits(limits); + let forwarded = self.forwarded.clone(); + let future = Box::pin(async move { + while let Some(frame) = agent.rx.next().await { + if let TransportFrame::Single(message) = frame.into_frame() { + forwarded.send(message).ok(); + } + } + Ok(()) + }); + (transport, future) + } + } + + #[tokio::test] + async fn oversized_websocket_frame_closes_live_agent_connection_and_registry() { + let (forwarded_tx, mut forwarded_rx) = mpsc::unbounded_channel(); + let registry = Arc::new(ConnectionRegistry::new(Arc::new(LimitedAgentFactory { + forwarded: forwarded_tx, + }))); + let app = Router::new().route( + "/acp", + get({ + let registry = registry.clone(); + move |ws: WebSocketUpgrade| { + let registry = registry.clone(); + async move { handle_ws_upgrade(registry, ws) } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let (mut socket, _) = connect_async(format!("ws://{addr}/acp")).await.unwrap(); + let ordinary = serde_json::json!({ + "jsonrpc": "2.0", "method": "test/valid", "params": {} + }); + socket + .send(ClientWsMessage::Text(ordinary.to_string().into())) + .await + .unwrap(); + timeout(Duration::from_secs(2), forwarded_rx.recv()) + .await + .unwrap() + .expect("live agent receives first message"); + let oversized = serde_json::json!({ + "jsonrpc": "2.0", "method": "test/oversized", + "params": { "payload": "x".repeat(512) } + }); + socket + .send(ClientWsMessage::Text(oversized.to_string().into())) + .await + .unwrap(); + timeout(Duration::from_secs(2), async { + loop { + if registry.len().await == 0 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("rejected frame must release registry entry"); + assert!( + timeout(Duration::from_secs(2), socket.next()).await.is_ok(), + "rejected frame must terminate the socket" + ); + server.abort(); + } + + #[tokio::test] + async fn saturated_websocket_frame_closes_stalled_agent_connection() { + let registry = Arc::new(ConnectionRegistry::new(Arc::new(StalledAgentFactory))); + let app = Router::new().route( + "/acp", + get({ + let registry = registry.clone(); + move |ws: WebSocketUpgrade| { + let registry = registry.clone(); + async move { handle_ws_upgrade(registry, ws) } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let (mut socket, _) = connect_async(format!("ws://{addr}/acp")).await.unwrap(); + let frame = json!({ + "jsonrpc": "2.0", "method": "test/stalled", "params": {"payload": "x".repeat(400)} + }) + .to_string(); + for _ in 0..4 { + // The agent deliberately does not consume its input. + if socket + .send(ClientWsMessage::Text(frame.clone().into())) + .await + .is_err() + { + break; + } + } + timeout(Duration::from_secs(2), async { + loop { + if registry.len().await == 0 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("saturated admission must release registry entry"); + assert!(timeout(Duration::from_secs(2), socket.next()).await.is_ok()); + server.abort(); + } + struct BatchAgentFactory { forwarded: mpsc::UnboundedSender>, } diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs index 3b399685..47300394 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs @@ -160,6 +160,20 @@ fn valid_origin(headers: &HeaderMap) -> bool { && origin == format!("http://{host}") } +/// HTTP qvalues are decimal 0..1 with at most three fractional digits, not +/// floating-point syntax (which also accepts NaN, exponents, and signs). +fn positive_quality(value: &str) -> Option { + let (whole, fraction) = value.split_once('.').unwrap_or((value, "")); + if fraction.len() > 3 || !fraction.bytes().all(|byte| byte.is_ascii_digit()) { + return None; + } + match whole { + "0" => Some(fraction.bytes().any(|byte| byte != b'0')), + "1" if fraction.bytes().all(|byte| byte == b'0') => Some(true), + _ => None, + } +} + fn accepts_both(headers: &HeaderMap) -> bool { let mut json = false; let mut sse = false; @@ -170,15 +184,20 @@ fn accepts_both(headers: &HeaderMap) -> bool { for item in value.split(',') { let mut parts = item.split(';'); let media = parts.next().unwrap_or("").trim(); - let mut quality = 1.0; + let mut quality = None; for part in parts { if let Some((key, q)) = part.trim().split_once('=') && key.trim().eq_ignore_ascii_case("q") { - quality = q.trim().parse::().unwrap_or(0.0); + let Some(positive) = positive_quality(q.trim()) else { + return false; + }; + if quality.replace(positive).is_some() { + return false; + } } } - if quality <= 0.0 || quality > 1.0 { + if quality == Some(false) { continue; } json |= media.eq_ignore_ascii_case("application/json"); @@ -756,6 +775,33 @@ mod tests { assert!(!accepts_both(&headers)); } + #[test] + fn accept_quality_uses_http_decimal_grammar() { + for quality in ["1", "1.", "1.000", "0.001", "0.5", "0.999"] { + let mut headers = HeaderMap::new(); + headers.insert( + "accept", + format!("application/json;q={quality}, text/event-stream") + .parse() + .unwrap(), + ); + assert!(accepts_both(&headers), "{quality}"); + } + for quality in [ + "NaN", "inf", "-1", "+1", "1e0", "0.0001", "1.001", "2", "", ".5", "00.5", "0", + "0.000", "1;q=0.9", + ] { + let mut headers = HeaderMap::new(); + headers.insert( + "accept", + format!("application/json;q={quality}, text/event-stream") + .parse() + .unwrap(), + ); + assert!(!accepts_both(&headers), "{quality}"); + } + } + #[test] fn mirrored_names_decode_canonical_base64() { assert!(matches_mirror( diff --git a/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs b/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs index fd3d5ac0..5adfcdc3 100644 --- a/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs +++ b/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs @@ -9,7 +9,7 @@ use std::{ }; use agent_client_protocol::{ - Agent, Client, Error, Responder, V2ConnectionTo, + Agent, Channel, Client, Error, Responder, V2ConnectionTo, mcp_server::McpServer, schema::{ProtocolVersion, v2}, }; @@ -409,3 +409,108 @@ async fn native_acp_stateless_rmcp_lifecycle() -> Result<(), Error> { .await .expect("native ACP/rmcp operation or cleanup timed out") } + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn active_rmcp_handler_drops_before_clean_acp_eof_completes() -> Result<(), Error> { + tokio::time::timeout(Duration::from_secs(10), async { + let (handler_started_tx, handler_started_rx) = oneshot::channel(); + let (handler_dropped_tx, mut handler_dropped_rx) = oneshot::channel(); + let pending = Arc::new(Mutex::new(HashMap::from([( + "eof".to_owned(), + (handler_started_tx, handler_dropped_tx), + )]))); + let (peer_stop_tx, peer_stop_rx) = oneshot::channel::<()>(); + let (peer, client) = Channel::duplex(); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("rmcp-eof-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + ), + ), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected native rmcp declaration"); + }; + let server_id = server.server_id.clone(); + let call_cx = cx.clone(); + cx.spawn(async move { + let mut params = json!({"name": "hang", "arguments": {"probe": "eof"}}); + params["_meta"] = meta("eof"); + let _result = call_cx + .send_request( + v2::MessageMcpRequest::new(server_id, "rmcp-eof", "tools/call") + .params(params.as_object().unwrap().clone()), + ) + .block_task() + .await; + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new("rmcp-eof-session")) + }, + agent_client_protocol::on_receive_request!(), + ); + let peer_task = tokio::spawn(agent.connect_with(peer, async move |_cx| { + let _ = peer_stop_rx.await; + Ok(()) + })); + let client_task = tokio::spawn(Client.v2().connect_with(client, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("rmcp-eof-client", "1"), + )) + .block_task() + .await?; + let pending = pending.clone(); + let server = McpServer::::from_rmcp("rmcp-eof", move || { + let (subscription_started, _) = oneshot::channel(); + let (subscription_stopped, _) = oneshot::channel(); + Service { + _drop: DropSignal(Arc::new(Mutex::new(None))), + started: Arc::new(Mutex::new(Some(subscription_started))), + stopped: Arc::new(Mutex::new(Some(subscription_stopped))), + pending: pending.clone(), + } + }); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)? + .start_session() + .block_task() + .await?; + cx.incoming_closed().await; + Ok(()) + })); + + handler_started_rx + .await + .map_err(Error::into_internal_error)?; + let _ = peer_stop_tx.send(()); + peer_task.await.map_err(Error::into_internal_error)??; + client_task.await.map_err(Error::into_internal_error)??; + assert!( + matches!(handler_dropped_rx.try_recv(), Ok(())), + "rmcp handler future must be destroyed before ACP driver completion" + ); + Ok(()) + }) + .await + .expect("rmcp handler close/join on ACP EOF timed out") +} diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 8e22600a..f71bbf63 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -99,8 +99,9 @@ pub struct ConnectionLimits { /// messages. One maximum frame's worth is reserved for responses/cancellation. pub max_queued_bytes: usize, /// Per-queue item limit and runtime admission limit for pending requests, - /// running tasks, dynamic handlers, and deferred dispatch. Values below one - /// are treated as one. Byte capacity is enforced separately. + /// total live tasks (running plus waiting), dynamic handlers, and deferred + /// dispatch. Values below one are treated as one. Byte capacity is enforced + /// separately. pub max_queued_frames: usize, } @@ -1979,7 +1980,7 @@ impl< OutgoingMessage::is_control, OutgoingMessage::is_urgent, ); - let (new_task_tx, new_task_rx) = admission::channel_with_capacity(limits.max_queued_frames); + let (new_task_tx, new_task_rx) = task_actor::task_channel(limits.max_queued_frames); let (dynamic_handler_tx, dynamic_handler_rx) = admission::channel_with_capacity(limits.max_queued_frames); let pending_replies = PendingReplies::with_capacity(limits.max_queued_frames); @@ -2004,11 +2005,11 @@ impl< pending_replies.registrar(), protocol_mode, ); - let spawn_result = connection.spawn(async move { + let transport_driver = async move { let result = transport_future.await; drop(transport_completion_tx.send(result.clone())); result - }); + }; // Destructure the channel endpoints let Channel { @@ -2021,8 +2022,6 @@ impl< let future = crate::util::instrument_with_connection_name(name, { let connection = connection.clone(); async move { - let () = spawn_result?; - let background = async { let incoming = incoming_actor::incoming_protocol_actor( me.counterpart(), @@ -2033,25 +2032,13 @@ impl< incoming_actor::IncomingHandlers::new(handler, on_close), protocol_compat.clone(), ); - let other_actors = async { - futures::try_join!( - // Protocol layer: OutgoingMessage -> RawJsonRpcMessage - outgoing_actor::outgoing_protocol_actor( - outgoing_rx, - pending_replies, - transport_outgoing_tx, - protocol_compat, - connection.incoming_closed.clone(), - ), - task_actor::task_actor( - new_task_rx, - &connection, - limits.max_queued_frames - ), - runner.run_with_connection_to(connection.clone()), - )?; - Ok(()) - }; + let other_actors = outgoing_actor::outgoing_protocol_actor( + outgoing_rx, + pending_replies, + transport_outgoing_tx, + protocol_compat, + connection.incoming_closed.clone(), + ); // EOF can wake a pending request consumer, which may make // the task actor fail while close callbacks are running. @@ -2065,7 +2052,7 @@ impl< .await }; - run_until_connection_close( + let lifecycle = run_until_connection_close( async { let result = background.await; connection.incoming_closed.request_shutdown(); @@ -2077,6 +2064,33 @@ impl< result }, connection.incoming_closed.clone(), + ); + // Only native operation supervisors are joined at shutdown. + // Ordinary spawned tasks and user runners remain disposable. + crate::util::run_until( + finish_actor_error( + runner.run_with_connection_to(connection.clone()), + &connection, + ), + crate::util::run_until( + finish_actor_error(transport_driver, &connection), + crate::util::run_until( + finish_actor_error( + task_actor::task_actor( + new_task_rx, + &connection, + limits.max_queued_frames, + ), + &connection, + ), + async { + let result = lifecycle.await; + connection.incoming_closed.request_shutdown(); + connection.wait_protected_operations().await; + result + }, + ), + ), ) .await } @@ -2086,6 +2100,23 @@ impl< } } +/// An EOF may wake a task that reports an error while on-close callbacks are +/// still running. Preserve that error without dropping the callback future. +async fn finish_actor_error( + future: impl Future>, + connection: &ConnectionTo, +) -> Result<(), crate::Error> { + let result = future.await; + if result.is_err() { + connection.incoming_closed.request_shutdown(); + if connection.incoming_closed.is_closing() { + connection.incoming_closed.closed().await; + } + connection.wait_protected_operations().await; + } + result +} + #[cfg(feature = "unstable_mcp_over_acp")] impl< Host: Role, @@ -3831,10 +3862,26 @@ pub struct ConnectionTo { )] protocol_mode: ProtocolMode, incoming_closed: IncomingClosed, + protected_operations: Arc>, } type SharedTransportCompletion = future::Shared>>; +#[derive(Default)] +struct ProtectedOperations { + pending: Vec>, + joining: Option>>, +} + +impl Debug for ProtectedOperations { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProtectedOperations") + .field("pending", &self.pending.len()) + .finish_non_exhaustive() + } +} + #[derive(Clone)] struct IncomingClosed { state: Arc, @@ -4009,7 +4056,65 @@ impl ConnectionTo { pending_replies, protocol_mode, incoming_closed: IncomingClosed::new(), + protected_operations: Arc::default(), + } + } + + #[cfg(feature = "unstable_mcp_over_acp")] + #[track_caller] + pub(crate) fn spawn_protected( + &self, + task: impl IntoFuture, IntoFuture: Send + 'static>, + ) -> Result<(), crate::Error> { + let (done_tx, done_rx) = oneshot::channel(); + let task = task.into_future(); + let mut state = self + .protected_operations + .lock() + .expect("protected operations poisoned"); + if state.joining.is_some() { + return Err(crate::Error::request_cancelled()); } + // Completed acknowledgments must not accumulate for the connection's + // entire lifetime. With completed entries reaped at each admission, + // this registry is bounded by the shared live-task limit. + state + .pending + .retain_mut(|done| matches!(done.try_recv(), Ok(None))); + self.spawn(async move { + let result = task.await; + let _ = done_tx.send(()); + result + })?; + state.pending.push(done_rx); + Ok(()) + } + + pub(crate) async fn wait_protected_operations(&self) { + let joining = { + let mut state = self + .protected_operations + .lock() + .expect("protected operations poisoned"); + if state.joining.is_none() { + let operations = std::mem::take(&mut state.pending); + state.joining = Some( + async move { + for operation in operations { + let _ = operation.await; + } + } + .boxed() + .shared(), + ); + } + state.joining.as_ref().expect("join initialized").clone() + }; + joining.await; + } + + pub(crate) fn request_shutdown(&self) { + self.incoming_closed.request_shutdown(); } #[cfg(feature = "unstable_protocol_v2")] @@ -7017,6 +7122,44 @@ impl FramePermit { .sum::() } + /// Reserve separately measured metadata retained after consuming this frame. + /// + /// This reserves `bytes` in every distinct budget covering the source frame, + /// including imported frames. It does not share or release the payload's + /// charge. Retain the returned permit with the metadata, then drop the + /// source permit once its payload has been consumed. + /// + /// Metadata uses data capacity, never the response/cancellation reserve. + /// Failure is immediate and releases any partial reservations; waiting + /// here could deadlock on capacity held by the source frame itself. + pub fn try_reserve_metadata(&self, bytes: usize) -> Result { + fn collect_budgets<'a>(permit: &'a FramePermit, budgets: &mut Vec<&'a Arc>) { + if !budgets + .iter() + .any(|budget| Arc::ptr_eq(budget, &permit.inner.budget)) + { + budgets.push(&permit.inner.budget); + } + for additional in &permit.additional { + collect_budgets(additional, budgets); + } + } + + let reserve = |budget: &Arc| { + budget.try_reserve(bytes, true).ok_or_else(|| { + crate::Error::invalid_request().data("retained metadata byte capacity exceeded") + }) + }; + let mut budgets = Vec::new(); + collect_budgets(self, &mut budgets); + let mut budgets = budgets.into_iter(); + let mut permit = reserve(budgets.next().expect("source frame has a budget"))?; + for budget in budgets { + permit.join(reserve(budget)?); + } + Ok(permit) + } + fn join(&mut self, other: FramePermit) { self.additional.push(other); } @@ -7291,7 +7434,27 @@ impl BudgetedFrame { /// Pollable receive half of an in-memory duplex. #[derive(Debug)] -pub struct FrameReceiver(std::pin::Pin>>); +pub struct FrameReceiver { + rx: std::pin::Pin>>, + _slots: async_channel::Sender<()>, +} + +/// A slot is reserved before a frame enters the queue and returned at dequeue. +/// Its drop also returns reservations abandoned by a cancelled send or sink. +#[derive(Debug)] +struct FrameSlot(async_channel::Sender<()>); + +impl Drop for FrameSlot { + fn drop(&mut self) { + let _ = self.0.try_send(()); + } +} + +#[derive(Debug)] +struct QueuedFrame { + frame: BudgetedFrame, + slot: FrameSlot, +} impl futures::Stream for FrameReceiver { type Item = BudgetedFrame; @@ -7300,16 +7463,24 @@ impl futures::Stream for FrameReceiver { mut self: std::pin::Pin<&mut Self>, cx: &mut Context<'_>, ) -> Poll> { - self.0.as_mut().poll_next(cx) + self.rx.as_mut().poll_next(cx).map(|item| { + item.map(|queued| { + drop(queued.slot); + queued.frame + }) + }) } } /// Backpressured frame sink. A synchronous send fails when its finite queue is full; /// asynchronous producers should use [`SinkExt::send`] instead. pub struct FrameSender { - tx: async_channel::Sender, + tx: async_channel::Sender, + slots: Box>, + slot_return: async_channel::Sender<()>, budget: Arc, - pending: Mutex>>>, + ready: Option, + waiting: Mutex>>>, } impl std::fmt::Debug for FrameSender { @@ -7371,8 +7542,11 @@ impl Clone for FrameSender { fn clone(&self) -> Self { Self { tx: self.tx.clone(), + slots: Box::new((*self.slots).clone()), + slot_return: self.slot_return.clone(), budget: self.budget.clone(), - pending: Mutex::new(None), + ready: None, + waiting: Mutex::new(None), } } } @@ -7401,6 +7575,47 @@ impl FrameSendError { } impl FrameSender { + async fn wait_for_slot( + tx: async_channel::Sender, + slots: async_channel::Receiver<()>, + slot_return: async_channel::Sender<()>, + ) -> Result { + if tx.is_closed() { + return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); + } + match future::select(Box::pin(slots.recv()), Box::pin(tx.closed())).await { + Either::Left((Ok(()), _)) if !tx.is_closed() => Ok(FrameSlot(slot_return)), + _ => Err(crate::Error::invalid_request().data("outgoing frame queue closed")), + } + } + + fn try_slot(&self) -> Result { + if self.tx.is_closed() { + return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); + } + self.slots + .try_recv() + .map(|()| FrameSlot(self.slot_return.clone())) + .map_err(|_| { + crate::Error::invalid_request().data("outgoing frame queue full or closed") + }) + } + + async fn reserve_slot(&self) -> Result { + Self::wait_for_slot( + self.tx.clone(), + (*self.slots).clone(), + self.slot_return.clone(), + ) + .await + } + + fn enqueue(&self, frame: BudgetedFrame, slot: FrameSlot) -> Result<(), crate::Error> { + self.tx + .try_send(QueuedFrame { frame, slot }) + .map_err(crate::util::internal_error) + } + /// Obtain the byte admission handle without retaining this channel sender. pub fn admission(&self) -> FrameAdmission { FrameAdmission(self.budget.clone()) @@ -7408,11 +7623,22 @@ impl FrameSender { /// Fail immediately rather than blocking a protocol dispatcher on its own output. pub fn try_send(&self, frame: TransportFrame) -> Result<(), FrameSendError> { + let Ok(slot) = self.try_slot() else { + return Err(FrameSendError { + frame: Box::new(frame), + reason: "outgoing frame queue full or closed", + }); + }; let budgeted = self.admission().try_admit(frame)?; - self.tx.try_send(budgeted).map_err(|error| FrameSendError { - frame: Box::new(error.into_inner().frame), - reason: "outgoing frame queue full or closed", - }) + self.tx + .try_send(QueuedFrame { + frame: budgeted, + slot, + }) + .map_err(|error| FrameSendError { + frame: Box::new(error.into_inner().frame.frame), + reason: "outgoing frame queue full or closed", + }) } /// Await byte and frame capacity outside ordered dispatch. @@ -7429,10 +7655,8 @@ impl FrameSender { return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); } }; - self.tx - .send(BudgetedFrame { frame, permit }) - .await - .map_err(crate::util::internal_error) + let slot = self.reserve_slot().await?; + self.enqueue(BudgetedFrame { frame, permit }, slot) } /// Transfer an application message's charge into its framed representation. @@ -7445,16 +7669,13 @@ impl FrameSender { ) -> Result<(), crate::Error> { let bytes = frame.to_json()?.len(); permit.cover_frame(bytes, !frame.is_control())?; - self.tx - .send(BudgetedFrame { frame, permit }) - .await - .map_err(crate::util::internal_error) + let slot = self.reserve_slot().await?; + self.enqueue(BudgetedFrame { frame, permit }, slot) } /// Stop accepting frames on this queue. pub fn close_channel(&self) { self.tx.close(); - self.pending.lock().expect("frame sender poisoned").take(); } /// Return whether the receiving endpoint has closed. @@ -7470,7 +7691,41 @@ impl Sink for FrameSender { self: std::pin::Pin<&mut Self>, cx: &mut std::task::Context<'_>, ) -> std::task::Poll> { - self.poll_flush(cx) + let this = self.get_mut(); + if this.tx.is_closed() { + this.ready.take(); + this.waiting + .get_mut() + .expect("frame sender poisoned") + .take(); + return Poll::Ready(Err( + crate::Error::invalid_request().data("outgoing frame queue closed") + )); + } + if this.ready.is_some() { + return Poll::Ready(Ok(())); + } + let waiting = this.waiting.get_mut().expect("frame sender poisoned"); + if waiting.is_none() { + *waiting = Some(Box::pin(Self::wait_for_slot( + this.tx.clone(), + (*this.slots).clone(), + this.slot_return.clone(), + ))); + } + match waiting.as_mut().expect("slot waiter").as_mut().poll(cx) { + Poll::Pending => Poll::Pending, + Poll::Ready(result) => { + waiting.take(); + match result { + Ok(slot) => { + this.ready = Some(slot); + Poll::Ready(Ok(())) + } + Err(error) => Poll::Ready(Err(error)), + } + } + } } fn start_send( @@ -7478,6 +7733,10 @@ impl Sink for FrameSender { mut item: BudgetedFrame, ) -> Result<(), Self::Error> { let this = self.get_mut(); + let slot = this + .ready + .take() + .ok_or_else(|| crate::Error::invalid_request().data("frame sender not ready"))?; if this.tx.is_closed() { return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); } @@ -7496,37 +7755,19 @@ impl Sink for FrameSender { })?; item.permit.join(permit); } - let mut pending = this.pending.lock().expect("frame sender poisoned"); - if pending.is_some() { - return Err(crate::Error::invalid_request().data("frame sender not ready")); - } - let tx = this.tx.clone(); - *pending = Some(Box::pin(async move { - tx.send(item).await.map_err(crate::util::internal_error) - })); - Ok(()) + this.enqueue(item, slot) } fn poll_flush( self: std::pin::Pin<&mut Self>, - cx: &mut std::task::Context<'_>, + _cx: &mut std::task::Context<'_>, ) -> std::task::Poll> { let this = self.get_mut(); - let mut pending = this.pending.lock().expect("frame sender poisoned"); - let Some(send) = pending.as_mut() else { - return Poll::Ready(if this.tx.is_closed() { - Err(crate::Error::invalid_request().data("outgoing frame queue closed")) - } else { - Ok(()) - }); - }; - match send.as_mut().poll(cx) { - Poll::Pending => Poll::Pending, - Poll::Ready(result) => { - pending.take(); - Poll::Ready(result) - } - } + Poll::Ready(if this.tx.is_closed() { + Err(crate::Error::invalid_request().data("outgoing frame queue closed")) + } else { + Ok(()) + }) } fn poll_close( @@ -7562,21 +7803,39 @@ impl Channel { }); let (a_tx, b_rx) = async_channel::bounded(limits.max_queued_frames.max(1)); let (b_tx, a_rx) = async_channel::bounded(limits.max_queued_frames.max(1)); + let (a_slot_tx, a_slots) = async_channel::bounded(limits.max_queued_frames.max(1)); + let (b_slot_tx, b_slots) = async_channel::bounded(limits.max_queued_frames.max(1)); + for _ in 0..limits.max_queued_frames.max(1) { + a_slot_tx.try_send(()).expect("initial queue slot"); + b_slot_tx.try_send(()).expect("initial queue slot"); + } ( Self { - rx: FrameReceiver(Box::pin(a_rx)), + rx: FrameReceiver { + rx: Box::pin(a_rx), + _slots: b_slot_tx.clone(), + }, tx: FrameSender { tx: a_tx, + slots: Box::new(a_slots), + slot_return: a_slot_tx.clone(), budget: budget.clone(), - pending: Mutex::new(None), + ready: None, + waiting: Mutex::new(None), }, }, Self { - rx: FrameReceiver(Box::pin(b_rx)), + rx: FrameReceiver { + rx: Box::pin(b_rx), + _slots: a_slot_tx.clone(), + }, tx: FrameSender { tx: b_tx, + slots: Box::new(b_slots), + slot_return: b_slot_tx, budget, - pending: Mutex::new(None), + ready: None, + waiting: Mutex::new(None), }, }, ) @@ -7963,28 +8222,118 @@ mod tests { } #[test] - fn closing_sender_drops_pending_sink_frame_charge() { + fn sink_clones_share_slots_and_dequeue_restores_capacity() { let frame = capacity_frame(); let bytes = frame.to_json().unwrap().len(); - let (mut left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + let (left, mut right) = Channel::duplex_with_limits(ConnectionLimits { max_frame_bytes: bytes, - max_queued_bytes: bytes * 4, + max_queued_bytes: bytes * 100, max_queued_frames: 1, }); - left.tx.try_send(frame.clone()).unwrap(); - let admitted = left.tx.admission().try_admit(frame).unwrap(); - std::pin::Pin::new(&mut left.tx) - .start_send(admitted) + let mut clones = (0..32).map(|_| left.tx.clone()).collect::>(); + clones[0] + .feed(left.tx.admission().try_admit(frame.clone()).unwrap()) + .now_or_never() + .expect("first feed must complete") .unwrap(); + for tx in &mut clones[1..] { + let admitted = left.tx.admission().try_admit(frame.clone()).unwrap(); + assert!(tx.feed(admitted).now_or_never().is_none()); + } + assert!(left.tx.try_send(frame.clone()).is_err()); + drop(right.rx.next().now_or_never().unwrap()); + clones[31] + .feed(left.tx.admission().try_admit(frame).unwrap()) + .now_or_never() + .expect("dequeue restores a slot") + .unwrap(); + } + + #[test] + fn dropping_ready_sink_releases_reserved_slot() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, _right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 8, + max_queued_frames: 1, + }); + let mut reserved = left.tx.clone(); assert!( - std::pin::Pin::new(&mut left.tx) - .poll_flush(&mut Context::from_waker(futures::task::noop_waker_ref())) - .is_pending() + std::pin::Pin::new(&mut reserved) + .poll_ready(&mut Context::from_waker(futures::task::noop_waker_ref())) + .is_ready() ); - left.tx.close_channel(); - assert_eq!(left.tx.budget.state.lock().unwrap().used, bytes); - drop(right.rx.next().now_or_never().unwrap()); - assert_eq!(left.tx.budget.state.lock().unwrap().used, 0); + assert!(left.tx.try_send(frame.clone()).is_err()); + drop(reserved); + left.tx + .try_send(frame) + .expect("dropping readiness frees slot"); + } + + #[tokio::test] + async fn closing_receiver_wakes_slot_blocked_sink_and_sender() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 8, + max_queued_frames: 1, + }); + left.tx.try_send(frame.clone()).unwrap(); + let mut waiting_sink = left.tx.clone(); + let mut waiting_feed = + Box::pin(waiting_sink.feed(left.tx.admission().try_admit(frame.clone()).unwrap())); + assert!(waiting_feed.as_mut().now_or_never().is_none()); + let mut waiting_send = Box::pin(left.tx.send_frame(frame)); + assert!(waiting_send.as_mut().now_or_never().is_none()); + drop(right); + assert!(waiting_feed.await.is_err()); + assert!(waiting_send.await.is_err()); + } + + #[test] + fn retained_metadata_charges_each_source_budget_without_pinning_payload() { + let (source, _source_peer) = Channel::duplex(); + let (destination, _destination_peer) = Channel::duplex(); + let mut payload = source.tx.budget.try_reserve(512, true).unwrap(); + payload.join(payload.clone()); + payload.join(destination.tx.budget.try_reserve(512, true).unwrap()); + + let metadata = payload.try_reserve_metadata(8).unwrap(); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 520); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 520); + drop(payload); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 8); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 8); + + let copy = metadata.clone(); + drop(metadata); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 8); + drop(copy); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 0); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 0); + } + + #[test] + fn retained_metadata_failure_rolls_back_and_cannot_use_control_capacity() { + let (source, _source_peer) = Channel::duplex(); + let (destination, _destination_peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 64, + max_queued_bytes: 128, + max_queued_frames: 4, + }); + let mut payload = source.tx.budget.try_reserve(60, true).unwrap(); + payload.join(destination.tx.budget.try_reserve(60, true).unwrap()); + assert!(payload.try_reserve_metadata(8).is_err()); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 60); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 60); + + let metadata = payload.try_reserve_metadata(4).unwrap(); + assert!(payload.try_reserve_metadata(1).is_err()); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 64); + drop(metadata); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 60); } fn application_channel( @@ -8128,7 +8477,7 @@ mod tests { }); let admission = channel.tx.admission(); let (message_tx, _message_rx) = application_channel(admission.clone()); - let (task_tx, mut task_rx) = admission::channel(); + let (task_tx, mut task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let pending = PendingReplies::default(); let connection = ConnectionTo::new( @@ -8303,15 +8652,72 @@ mod tests { assert!(rx.next().now_or_never().unwrap().unwrap().is_urgent()); } + #[test] + fn cancellation_admission_is_urgent_at_every_queue_occupancy() { + use super::admission::ReceiverClose as _; + + for queued in 0..=3 { + for asynchronous in [false, true] { + for wrapped in [false, true] { + let (channel, _peer) = Channel::duplex_with_limits(ConnectionLimits { + max_queued_frames: 3, + ..ConnectionLimits::default() + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + for _ in 0..queued { + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("data", serde_json::json!({})).unwrap(), + }) + .unwrap(); + } + let params = serde_json::json!({"requestId":"waiting"}); + let cancel = if wrapped { + UntypedMessage::new( + "_proxy/successor", + serde_json::json!({"method":"$/cancel_request","params":params}), + ) + } else { + UntypedMessage::new("$/cancel_request", params) + } + .unwrap(); + let message = OutgoingMessage::Notification { untyped: cancel }; + if asynchronous { + tx.send(message) + .now_or_never() + .expect("urgent admission does not wait for ordinary queue space") + .unwrap(); + } else { + tx.unbounded_send(message).unwrap(); + } + let urgent = future::poll_fn(|cx| rx.poll_urgent(cx)) + .now_or_never() + .expect("readiness gate must observe cancellation") + .unwrap(); + assert!(urgent.is_urgent()); + for _ in 0..queued { + assert!(!rx.next().now_or_never().unwrap().unwrap().is_urgent()); + } + assert!(rx.next().now_or_never().is_none()); + } + } + } + } + #[test] fn cancellation_passes_waiting_request_and_saturated_data_lane() { + for capacity in [1, 3, ConnectionLimits::default().max_queued_frames] { + check_cancellation_passes_waiting_request(capacity); + } + } + + fn check_cancellation_passes_waiting_request(capacity: usize) { let (transport, mut peer) = Channel::duplex_with_limits(ConnectionLimits { max_frame_bytes: 1024, max_queued_bytes: 4096, - max_queued_frames: 1, + max_queued_frames: capacity, }); let (message_tx, message_rx) = application_channel(transport.tx.admission()); - let (task_tx, _task_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_tx, _dynamic_rx) = admission::channel(); let pending_replies = PendingReplies::default(); let connection = ConnectionTo::new( @@ -8369,13 +8775,19 @@ mod tests { #[test] fn cancellation_of_queued_request_does_not_wait_for_unrelated_readiness() { + for capacity in [1, 3, ConnectionLimits::default().max_queued_frames] { + check_cancellation_of_queued_request(capacity); + } + } + + fn check_cancellation_of_queued_request(capacity: usize) { let (transport, mut peer) = Channel::duplex_with_limits(ConnectionLimits { max_frame_bytes: 1024, max_queued_bytes: 8192, - max_queued_frames: 1, + max_queued_frames: capacity, }); let (message_tx, message_rx) = application_channel(transport.tx.admission()); - let (task_tx, _task_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_tx, _dynamic_rx) = admission::channel(); let pending_replies = PendingReplies::default(); let connection = ConnectionTo::new( @@ -8430,13 +8842,61 @@ mod tests { first.detach(); } + #[cfg(feature = "unstable_mcp_over_acp")] + #[test] + fn protected_operation_tracking_is_bounded_across_sequential_completions() { + let (message_tx, _message_rx) = admission::channel(); + let (task_tx, task_rx) = task_actor::task_channel(2); + let (dynamic_tx, _dynamic_rx) = admission::channel(); + let pending_replies = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx, + task_tx, + dynamic_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending_replies.registrar(), + ProtocolMode::disabled(), + ); + let mut driver = Box::pin(task_actor::task_actor(task_rx, &connection, 2)); + for _ in 0..1000 { + connection.spawn_protected(async { Ok(()) }).unwrap(); + assert_eq!( + connection + .protected_operations + .lock() + .unwrap() + .pending + .len(), + 1, + "previous completions must be reaped before another admission" + ); + assert!(driver.as_mut().now_or_never().is_none()); + } + assert!( + connection + .wait_protected_operations() + .now_or_never() + .is_some() + ); + assert!( + connection + .protected_operations + .lock() + .unwrap() + .pending + .is_empty() + ); + assert!(connection.spawn_protected(async { Ok(()) }).is_err()); + } + #[cfg(feature = "unstable_protocol_v2")] fn connection_with_task_receiver() -> ( ConnectionTo, admission::SimpleReceiver, ) { let (message_tx, _message_rx) = admission::channel(); - let (task_tx, task_rx) = admission::channel(); + let (task_tx, task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); @@ -8561,7 +9021,7 @@ mod tests { #[test] fn v2_proxy_rejects_explicitly_prewrapped_initialize_request() { let (message_tx, message_rx) = admission::channel(); - let (task_tx, _task_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); @@ -8723,7 +9183,7 @@ mod tests { admission::SimpleReceiver>, ) { let (message_tx, _message_rx) = admission::channel(); - let (task_tx, _task_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_handler_tx, dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); @@ -8765,7 +9225,7 @@ mod tests { PendingReplies, ) { let (message_tx, message_rx) = admission::channel(); - let (task_tx, _task_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); @@ -8797,7 +9257,7 @@ mod tests { let (channel, _) = Channel::duplex_with_limits(limits); let admission = channel.tx.admission(); let (message_tx, message_rx) = application_channel(admission.clone()); - let (task_tx, _task_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let pending = PendingReplies::default(); let connection = ConnectionTo::new( @@ -9248,7 +9708,7 @@ mod tests { #[test] fn ordered_request_is_marked_before_entering_outgoing_queue() { let (message_tx, mut message_rx) = admission::channel(); - let (task_tx, mut task_rx) = admission::channel(); + let (task_tx, mut task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); @@ -9335,7 +9795,7 @@ mod tests { #[test] fn v2_dynamic_handler_guard_registers_and_removes_handler() { let (message_tx, _message_rx) = admission::channel(); - let (task_tx, _task_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); let (dynamic_handler_tx, mut dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); diff --git a/src/agent-client-protocol/src/jsonrpc/admission.rs b/src/agent-client-protocol/src/jsonrpc/admission.rs index 8bb7f1ce..4c5d8a1d 100644 --- a/src/agent-client-protocol/src/jsonrpc/admission.rs +++ b/src/agent-client-protocol/src/jsonrpc/admission.rs @@ -57,12 +57,14 @@ impl Sender { } pub fn unbounded_send(&self, item: T) -> Result<(), SendError> { - // Preserve the order of a queued request followed by its cancellation - // when the ordinary lane has room. Use the bypass lane if ordinary - // capacity is exhausted or its consumer is blocked on readiness. - let urgent = self.inner.admission.as_ref().is_some_and(|admission| { - (admission.urgent)(&item) && (self.inner.tx.is_empty() || self.inner.tx.is_full()) - }); + // A readiness-blocked consumer polls only the urgent lane, regardless + // of ordinary queue occupancy. The outgoing actor settles cancellation + // locally if its request has not yet been published. + let urgent = self + .inner + .admission + .as_ref() + .is_some_and(|admission| (admission.urgent)(&item)); let item = if let Some(admission) = &self.inner.admission { let bytes = (admission.measure)(&item).map_err(|error| SendError { item: None, @@ -110,9 +112,11 @@ impl Sender { } pub async fn send(&self, item: T) -> Result<(), crate::Error> { - let urgent = self.inner.admission.as_ref().is_some_and(|admission| { - (admission.urgent)(&item) && (self.inner.tx.is_empty() || self.inner.tx.is_full()) - }); + let urgent = self + .inner + .admission + .as_ref() + .is_some_and(|admission| (admission.urgent)(&item)); let tx = if urgent { self.inner.urgent_tx.as_ref().expect("urgent lane exists") } else { diff --git a/src/agent-client-protocol/src/jsonrpc/run.rs b/src/agent-client-protocol/src/jsonrpc/run.rs index 571bf61a..f5aaa870 100644 --- a/src/agent-client-protocol/src/jsonrpc/run.rs +++ b/src/agent-client-protocol/src/jsonrpc/run.rs @@ -7,6 +7,8 @@ use std::future::Future; use std::marker::PhantomData; +use futures::future::{Either, select}; + use crate::{ ConnectionTo, jsonrpc::{ConnectionContext, RawConnectionContext, connection_context}, @@ -67,8 +69,29 @@ where // Box the futures to avoid stack overflow with deeply nested RunIn chains let a_fut = Box::pin(self.a.run_with_connection_to(cx.clone())); let b_fut = Box::pin(self.b.run_with_connection_to(cx.clone())); - let ((), ()) = futures::future::try_join(a_fut, b_fut).await?; - Ok(()) + match select(a_fut, b_fut).await { + Either::Left((Ok(()), b)) => b.await, + Either::Right((Ok(()), a)) => a.await, + Either::Left((Err(error), b)) => { + cx.request_shutdown(); + // A different runner may own cleanup of a scoped MCP tool. + // Continue polling it without waiting for an unrelated + // never-ending runner after protected operations finish. + match select(b, Box::pin(cx.wait_protected_operations())).await { + Either::Left((_, cleanup)) => cleanup.await, + Either::Right(((), _)) => {} + } + Err(error) + } + Either::Right((Err(error), a)) => { + cx.request_shutdown(); + match select(a, Box::pin(cx.wait_protected_operations())).await { + Either::Left((_, cleanup)) => cleanup.await, + Either::Right(((), _)) => {} + } + Err(error) + } + } } } diff --git a/src/agent-client-protocol/src/jsonrpc/task_actor.rs b/src/agent-client-protocol/src/jsonrpc/task_actor.rs index cc30a9e7..8f4e8289 100644 --- a/src/agent-client-protocol/src/jsonrpc/task_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/task_actor.rs @@ -1,15 +1,48 @@ use std::panic::Location; +use std::sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, +}; +use futures::channel::oneshot; +use futures::future::{self, Either}; use futures::{FutureExt, StreamExt, future::BoxFuture}; use crate::ConnectionTo; use crate::role::Role; -pub type TaskTx = super::admission::Sender; +#[derive(Clone, Debug)] +pub struct TaskTx { + sender: super::admission::Sender, + live: Arc, + capacity: usize, +} + +pub fn task_channel(capacity: usize) -> (TaskTx, super::admission::SimpleReceiver) { + let capacity = capacity.max(1); + let (sender, receiver) = super::admission::channel_with_capacity(capacity); + ( + TaskTx { + sender, + live: Arc::new(AtomicUsize::new(0)), + capacity, + }, + receiver, + ) +} + +struct LiveTask(Arc); + +impl Drop for LiveTask { + fn drop(&mut self) { + self.0.fetch_sub(1, Ordering::AcqRel); + } +} #[must_use] pub(crate) struct Task { future: BoxFuture<'static, Result<(), crate::Error>>, + live: Option, } impl Task { @@ -34,12 +67,21 @@ impl Task { } }, ) - .boxed() + .boxed(), + live: None, } } - pub fn spawn(self, task_tx: &TaskTx) -> Result<(), crate::Error> { + pub fn spawn(mut self, task_tx: &TaskTx) -> Result<(), crate::Error> { + task_tx + .live + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |live| { + (live < task_tx.capacity).then_some(live + 1) + }) + .map_err(|_| crate::util::internal_error("live task capacity exceeded"))?; + self.live = Some(LiveTask(task_tx.live.clone())); task_tx + .sender .unbounded_send(self) .map_err(crate::util::internal_error)?; Ok(()) @@ -54,12 +96,79 @@ impl Task { /// The "task actor" manages dynamically spawned tasks. pub(super) async fn task_actor( task_rx: super::admission::SimpleReceiver, - _cx: &ConnectionTo, + cx: &ConnectionTo, max_running_tasks: usize, ) -> Result<(), crate::Error> { - use futures::TryStreamExt as _; - task_rx - .map(Ok::<_, crate::Error>) - .try_for_each_concurrent(max_running_tasks.max(1), |task| task.future) - .await + let (error_tx, error_rx) = oneshot::channel(); + let first_error = Arc::new(Mutex::new(Some(error_tx))); + let running = task_rx.for_each_concurrent(max_running_tasks.max(1), |task| { + let first_error = first_error.clone(); + async move { + let Task { future, live } = task; + let result = future.await; + drop(live); + if let Err(error) = result + && let Some(tx) = first_error + .lock() + .expect("task error mutex poisoned") + .take() + { + drop(tx.send(error)); + } + } + }); + let on_error = async { + let error = error_rx + .await + .expect("task driver dropped before completion"); + cx.incoming_closed.request_shutdown(); + // Keep polling the driver while native supervisors finish. A failed + // disposable task cannot drop those supervisors or force us to join + // arbitrary never-ending disposable tasks. + cx.wait_protected_operations().await; + Err(error) + }; + match future::select(Box::pin(running), Box::pin(on_error)).await { + Either::Left(((), _)) => Ok(()), + Either::Right((result, _)) => result, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use futures::channel::oneshot; + + #[test] + fn running_child_occupies_total_live_capacity_not_just_queue_capacity() { + futures::executor::block_on(async { + let (tx, mut rx) = task_channel(1); + let (child_done_tx, child_done_rx) = oneshot::channel::<()>(); + Task::new(Location::caller(), async move { + let _ = child_done_rx.await; + Ok(()) + }) + .spawn(&tx) + .unwrap(); + // The child is no longer in the waiting queue, but remains live. + let child = rx.next().await.unwrap(); + let (callback_dropped_tx, callback_dropped_rx) = oneshot::channel::<()>(); + let callback = async move { + let _drop_on_rejection = callback_dropped_tx; + futures::future::pending::<()>().await; + Ok(()) + }; + let rejection = Task::new(Location::caller(), callback) + .spawn(&tx) + .expect_err("an ordered callback cannot wait behind a permanent child"); + assert!(rejection.to_string().contains("live task capacity")); + assert!(callback_dropped_rx.now_or_never().unwrap().is_err()); + + drop(child_done_tx); + child.run_for_test().await.unwrap(); + Task::new(Location::caller(), async { Ok(()) }) + .spawn(&tx) + .expect("child completion releases total live capacity"); + }); + } } diff --git a/src/agent-client-protocol/src/mcp_server/active_session.rs b/src/agent-client-protocol/src/mcp_server/active_session.rs index 1df3fc3d..d0ed1b4f 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -378,7 +378,7 @@ where ); let is_discovery = method == "server/discover"; let shutdown_connection = connection.clone(); - connection.spawn(async move { + connection.spawn_protected(async move { let request = McpRequest { method, params }; let cleanup_connection = context.connection().clone(); let operation = service.execute(request, context); @@ -432,7 +432,7 @@ where // Keep the operation admitted until its backend has actually stopped. let (backend_stop_tx, backend_stop_rx) = oneshot::channel::<()>(); let (backend_done_tx, mut backend_done_rx) = oneshot::channel(); - let spawn_result = connection.spawn(async move { + let spawn_result = connection.spawn_protected(async move { // Own (not merely borrow) the future so cancellation drops its // backend before the completion acknowledgement is published. let run = Box::pin(backend.connect_to(server)); @@ -448,7 +448,7 @@ where responder.respond_with_error(error)?; return Ok(Handled::Yes); } - let spawn_result = connection.spawn(async move { + let spawn_result = connection.spawn_protected(async move { let inner_id = RequestId::Str(request_id.0.to_string()); let is_discovery = method == "server/discover"; let process = async { diff --git a/src/agent-client-protocol/src/mcp_server/context.rs b/src/agent-client-protocol/src/mcp_server/context.rs index 876390d4..cb2f3455 100644 --- a/src/agent-client-protocol/src/mcp_server/context.rs +++ b/src/agent-client-protocol/src/mcp_server/context.rs @@ -67,7 +67,7 @@ pub struct McpConnectionTo { } impl McpConnectionTo { - #[cfg(feature = "unstable_mcp_over_acp")] + #[cfg(all(feature = "unstable_mcp_over_acp", feature = "schemars"))] pub(crate) fn register_cleanup(&self, done: oneshot::Receiver<()>) { if let Some(cleanup) = &self.cleanup { cleanup.lock().expect("MCP cleanup poisoned").push(done); diff --git a/src/agent-client-protocol/tests/jsonrpc_request_cancellation.rs b/src/agent-client-protocol/tests/jsonrpc_request_cancellation.rs index 3be9a7c2..dd9fe176 100644 --- a/src/agent-client-protocol/tests/jsonrpc_request_cancellation.rs +++ b/src/agent-client-protocol/tests/jsonrpc_request_cancellation.rs @@ -2,11 +2,10 @@ //! //! These tests avoid sleeps by relying on two ordering guarantees: //! -//! - Messages are delivered in the order they were sent, and each side's -//! dispatch loop processes incoming messages sequentially. A request/response -//! round trip therefore acts as a barrier: by the time the response arrives, -//! every message sent before the request (including any `$/cancel_request`) -//! has been fully processed by the peer. +//! - Ordinary messages preserve queue order, and incoming dispatch is +//! sequential. A round trip after a request proves it reached the peer. +//! Cancellation has a separate urgent lane: canceling before publication +//! settles locally instead of sending either message to the peer. //! - Test handlers report observed cancellations through in-process channels, //! which the test awaits (with a timeout) instead of sleeping. @@ -53,6 +52,22 @@ async fn next_with_timeout(rx: &mut mpsc::UnboundedReceiver) -> T { .expect("channel closed before expected event") } +/// Remote-cancellation tests must first publish the request through every hop. +async fn publication_barrier(connection: &ConnectionTo) { + let response = tokio::time::timeout( + tokio::time::Duration::from_secs(10), + connection + .send_request(SimpleRequest { + message: "barrier".into(), + }) + .block_task(), + ) + .await + .expect("publication barrier timed out") + .expect("publication barrier failed"); + assert_eq!(response.result, "echo: barrier"); +} + /// Assert that no item is currently buffered on `rx`. /// /// Callers must first establish an ordering barrier (such as a @@ -432,6 +447,7 @@ async fn cancelling_request_sent_to_successor_peer_sends_wrapped_cancel() { .run_until(async { let (wrapped_cancel_tx, mut wrapped_cancel_rx) = mpsc::unbounded(); let (plain_cancel_tx, mut plain_cancel_rx) = mpsc::unbounded(); + let (started_tx, mut started_rx) = mpsc::unbounded(); let (server_reader, server_writer, client_reader, client_writer) = setup_test_streams(); let server_transport = @@ -440,9 +456,10 @@ async fn cancelling_request_sent_to_successor_peer_sends_wrapped_cancel() { .builder() .on_receive_request_from( WrappedSuccessor, - async |_request: SimpleRequest, - responder: Responder, - cx: ConnectionTo| { + async move |_request: SimpleRequest, + responder: Responder, + cx: ConnectionTo| { + started_tx.unbounded_send(responder.id().clone()).unwrap(); let cancellation = responder.cancellation(); cx.spawn(async move { let response = cancellation @@ -498,6 +515,7 @@ async fn cancelling_request_sent_to_successor_peer_sends_wrapped_cancel() { }, ); let expected_id = request.id().clone(); + assert_eq!(next_with_timeout(&mut started_rx).await, expected_id); request.cancel()?; let error = request .block_task() @@ -573,6 +591,7 @@ async fn sent_request_can_send_cancellation_for_its_id() { local .run_until(async { let (cancel_tx, mut cancel_rx) = mpsc::unbounded(); + let (started_tx, mut started_rx) = mpsc::unbounded(); let (server_reader, server_writer, client_reader, client_writer) = setup_test_streams(); let server_transport = @@ -580,9 +599,9 @@ async fn sent_request_can_send_cancellation_for_its_id() { let server = UntypedRole .builder() .on_receive_request( - async |request: SimpleRequest, - responder: Responder, - _connection: ConnectionTo| { + async move |request: SimpleRequest, + responder: Responder, + _connection: ConnectionTo| { if request.message == "barrier" { return responder.respond(SimpleResponse { result: format!("echo: {}", request.message), @@ -591,6 +610,7 @@ async fn sent_request_can_send_cancellation_for_its_id() { // Park other requests (by dropping the responder) so // the cancelled request is never answered and the // client handle stays unconsumed. + started_tx.unbounded_send(responder.id().clone()).unwrap(); Ok(()) }, agent_client_protocol::on_receive_request!(), @@ -619,6 +639,7 @@ async fn sent_request_can_send_cancellation_for_its_id() { message: "slow".into(), }); let expected_id = request.id().clone(); + assert_eq!(next_with_timeout(&mut started_rx).await, expected_id); request.cancel()?; let received = next_with_timeout(&mut cancel_rx).await; @@ -1367,6 +1388,7 @@ async fn forward_response_to_propagates_cancellation_to_downstream_request() { connection.send_request(SimpleRequest { message: "cancel downstream".into(), }); + publication_barrier(&connection).await; request.cancel()?; // The backend answers the parked request only once the @@ -1521,6 +1543,7 @@ async fn send_proxied_message_does_not_tunnel_cancel_notifications() { message: "park".into(), }); let client_request_id = request.id().clone(); + publication_barrier(&connection).await; request.cancel()?; let error = request @@ -1824,6 +1847,7 @@ async fn custom_forwarding_propagates_cancellation_when_opted_in() { message: "park".into(), }); let client_request_id = request.id().clone(); + publication_barrier(&connection).await; request.cancel()?; let error = request @@ -1905,6 +1929,7 @@ async fn custom_forwarding_absorbs_cancellation_by_default() { connection.send_request(SimpleRequest { message: "park".into(), }); + publication_barrier(&connection).await; request.cancel()?; // Barrier: the cancellation has now been processed by the diff --git a/src/agent-client-protocol/tests/live_task_capacity.rs b/src/agent-client-protocol/tests/live_task_capacity.rs new file mode 100644 index 00000000..51b69018 --- /dev/null +++ b/src/agent-client-protocol/tests/live_task_capacity.rs @@ -0,0 +1,80 @@ +use std::time::Duration; + +use agent_client_protocol::{ + BudgetedFrame, Channel, ConnectionLimits, Error, JsonRpcRequest, JsonRpcResponse, + RawJsonRpcMessage, TransportFrame, UntypedRole, +}; +use futures::{StreamExt as _, channel::oneshot}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)] +#[request(method = "_test/callback", response = CallbackResponse)] +struct CallbackRequest {} + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)] +struct CallbackResponse {} + +#[tokio::test] +async fn persistent_child_cannot_strand_an_ordered_response_callback() -> Result<(), Error> { + tokio::time::timeout(Duration::from_secs(10), async { + let (transport, mut peer) = Channel::duplex_with_limits(ConnectionLimits { + max_queued_frames: 1, + ..ConnectionLimits::default() + }); + let (published_tx, published_rx) = oneshot::channel(); + let (reply_tx, reply_rx) = oneshot::channel::<()>(); + let (child_transport, child_peer) = Channel::duplex(); + let connection = UntypedRole + .builder() + .connect_with(transport, async move |cx| { + let _child = cx.spawn_connection(UntypedRole.builder(), child_transport)?; + let (callback_tx, callback_rx) = oneshot::channel(); + let request = cx.send_request(CallbackRequest {}); + published_rx.await.map_err(Error::into_internal_error)?; + let error = request + .on_receiving_result(async move |_result| { + let _ = callback_tx.send(()); + Ok(()) + }) + .expect_err("the child reserves the sole live task slot"); + assert!(error.to_string().contains("live task capacity"), "{error}"); + assert!( + callback_rx.await.is_err(), + "rejected callback must release its captured resources" + ); + let _ = reply_tx.send(()); + cx.incoming_closed().await; + Ok(()) + }); + let reply_and_close = async move { + let request = loop { + match peer.rx.next().await.map(BudgetedFrame::into_frame) { + Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) => { + break request; + } + Some(TransportFrame::Single(RawJsonRpcMessage::Notification(_))) => {} + other => panic!("parent request did not reach the peer: {other:?}"), + } + }; + let _ = published_tx.send(()); + let _ = reply_rx.await; + peer.tx + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( + request.id, + Ok(serde_json::json!({})), + ))) + .await + .unwrap(); + // Half-close input but continue draining any cancellation/output. + // Dropping both directions here would make a legitimate late write + // fail for reasons unrelated to the response acknowledgment. + peer.tx.close_channel(); + while peer.rx.next().await.is_some() {} + }; + let (connection, ()) = futures::future::join(connection, reply_and_close).await; + drop(child_peer); + connection + }) + .await + .expect("persistent child stranded the response dispatcher") +} diff --git a/src/agent-client-protocol/tests/native_mcp_shutdown.rs b/src/agent-client-protocol/tests/native_mcp_shutdown.rs new file mode 100644 index 00000000..8b8fd7e6 --- /dev/null +++ b/src/agent-client-protocol/tests/native_mcp_shutdown.rs @@ -0,0 +1,220 @@ +#![cfg(all(feature = "unstable_protocol_v2", feature = "unstable_mcp_over_acp"))] + +use std::{sync::Mutex, time::Duration}; + +use agent_client_protocol::{ + Agent, Channel, Client, ConnectionTo, Error, Responder, RunWithConnectionTo, V2ConnectionTo, + mcp_server::{McpOutcome, McpRequest, McpRequestContext, McpServer, McpService}, + schema::{ProtocolVersion, v2}, +}; +use futures::future::BoxFuture; +use serde_json::json; +use tokio::sync::oneshot; + +struct OnDrop(Option>); + +impl Drop for OnDrop { + fn drop(&mut self) { + if let Some(done) = self.0.take() { + let _ = done.send(()); + } + } +} + +struct CleanupService { + started: Mutex>>, + cleanup_started: Mutex>>, + runner_woke: Mutex>>, + release: Mutex>>, + dropped: Mutex>>, +} + +struct ShutdownRunner(oneshot::Sender<()>); + +impl RunWithConnectionTo for ShutdownRunner { + async fn run_with_connection_to(self, cx: ConnectionTo) -> Result<(), Error> { + cx.shutdown_requested().await; + let _ = self.0.send(()); + std::future::pending().await + } +} + +impl McpService for CleanupService { + fn execute( + &self, + _request: McpRequest, + context: McpRequestContext, + ) -> BoxFuture<'static, Result> { + let started = self.started.lock().unwrap().take().unwrap(); + let cleanup_started = self.cleanup_started.lock().unwrap().take().unwrap(); + let runner_woke = self.runner_woke.lock().unwrap().take().unwrap(); + let release = self.release.lock().unwrap().take().unwrap(); + let dropped = self.dropped.lock().unwrap().take().unwrap(); + Box::pin(async move { + let _drop = OnDrop(Some(dropped)); + let _ = started.send(()); + context.operation_cancellation().cancelled().await; + runner_woke.await.map_err(Error::into_internal_error)?; + let _ = cleanup_started.send(()); + let _ = release.await; + Err(Error::request_cancelled()) + }) + } +} + +#[derive(Clone, Copy)] +enum Shutdown { + PeerEof, + Foreground, + UnrelatedTaskError, +} + +async fn shutdown_joins_native_cleanup(shutdown: Shutdown) -> Result<(), Error> { + tokio::time::timeout(Duration::from_secs(10), async move { + let (started_tx, started_rx) = oneshot::channel(); + let (cleanup_started_tx, cleanup_started_rx) = oneshot::channel(); + let (runner_woke_tx, runner_woke_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + let (dropped_tx, mut dropped_rx) = oneshot::channel(); + let (peer_stop_tx, peer_stop_rx) = oneshot::channel::<()>(); + let mut peer_stop_tx = Some(peer_stop_tx); + let (client_stop_tx, client_stop_rx) = oneshot::channel::<()>(); + let (peer, client) = Channel::duplex(); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx| { + responder.respond(v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("shutdown-agent", "1"), + )) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected native MCP server declaration"); + }; + let server_id = server.server_id.clone(); + let request_connection = cx.clone(); + cx.spawn(async move { + let request = + v2::MessageMcpRequest::new(server_id, "cleanup-probe", "tools/call") + .params( + json!({ + "name": "probe", + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } + }) + .as_object() + .unwrap() + .clone(), + ); + let _result = request_connection.send_request(request).block_task().await; + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new("shutdown-session")) + }, + agent_client_protocol::on_receive_request!(), + ); + let peer_task = tokio::spawn(agent.connect_with(peer, async move |_cx| { + let _ = peer_stop_rx.await; + Ok(()) + })); + let client_task = tokio::spawn(Client.v2().connect_with(client, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("shutdown-client", "1"), + )) + .block_task() + .await?; + let server = McpServer::new_service( + CleanupService { + started: Mutex::new(Some(started_tx)), + cleanup_started: Mutex::new(Some(cleanup_started_tx)), + runner_woke: Mutex::new(Some(runner_woke_rx)), + release: Mutex::new(Some(release_rx)), + dropped: Mutex::new(Some(dropped_tx)), + }, + "shutdown-test", + ShutdownRunner(runner_woke_tx), + ); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)? + .start_session() + .block_task() + .await?; + match shutdown { + Shutdown::PeerEof => cx.incoming_closed().await, + Shutdown::Foreground => { + let _ = client_stop_rx.await; + } + Shutdown::UnrelatedTaskError => { + let _ = client_stop_rx.await; + cx.spawn(async { Err(Error::internal_error().data("unrelated task failed")) })?; + std::future::pending::<()>().await; + } + } + Ok(()) + })); + + started_rx.await.map_err(Error::into_internal_error)?; + if matches!(shutdown, Shutdown::PeerEof) { + let _ = peer_stop_tx.take().unwrap().send(()); + } else { + let _ = client_stop_tx.send(()); + } + cleanup_started_rx + .await + .map_err(Error::into_internal_error)?; + assert!( + !client_task.is_finished(), + "connection discarded pending native cleanup" + ); + assert!(matches!( + dropped_rx.try_recv(), + Err(oneshot::error::TryRecvError::Empty) + )); + let _ = release_tx.send(()); + dropped_rx.await.map_err(Error::into_internal_error)?; + let result = client_task.await.map_err(Error::into_internal_error)?; + if matches!(shutdown, Shutdown::UnrelatedTaskError) { + let error = result.expect_err("task failure must remain the primary connection error"); + assert!( + error.to_string().contains("unrelated task failed"), + "{error}" + ); + } else { + result?; + } + if !matches!(shutdown, Shutdown::PeerEof) { + let _ = peer_stop_tx.take().unwrap().send(()); + } + peer_task.await.map_err(Error::into_internal_error)??; + Ok(()) + }) + .await + .expect("native cleanup shutdown timed out") +} + +#[tokio::test] +async fn native_cleanup_survives_clean_incoming_eof() -> Result<(), Error> { + shutdown_joins_native_cleanup(Shutdown::PeerEof).await +} + +#[tokio::test] +async fn native_cleanup_survives_foreground_return() -> Result<(), Error> { + shutdown_joins_native_cleanup(Shutdown::Foreground).await +} + +#[tokio::test] +async fn unrelated_task_error_waits_for_native_cleanup_and_preserves_error() -> Result<(), Error> { + shutdown_joins_native_cleanup(Shutdown::UnrelatedTaskError).await +} From 7db867d1c957809856ce621800c6ecd03ce95ee2 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Mon, 28 Sep 2026 13:18:16 +0200 Subject: [PATCH 7/8] fix(acp)!: preserve raw errors and finalize HTTP admission Keep raw JSON-RPC errors protocol-neutral so connector-backed MCP preserves explicit null and extension fields. Reserve HTTP queue capacity before publishing metadata and tear down idle agents after fatal router failure. Add permanent regressions for all three review findings. BREAKING CHANGE: RawJsonRpcMessage::Response now carries RawJsonRpcResponse with a boxed RawJsonRpcError instead of an ACP-specific schema response. Typed ACP request consumers retain the existing Error API. --- md/http-transport.md | 10 + md/migration-stateless-mcp.md | 13 + md/transport-architecture.md | 9 +- .../src/trace.rs | 5 +- src/agent-client-protocol-http/src/client.rs | 5 +- .../src/connection.rs | 180 +++++++--- .../src/connection_admission_tests.rs | 57 +++- .../src/http_server.rs | 73 +++- .../src/websocket_server.rs | 4 +- src/agent-client-protocol/CHANGELOG.md | 18 +- src/agent-client-protocol/src/jsonrpc.rs | 33 +- .../src/jsonrpc/incoming_actor.rs | 5 +- .../src/jsonrpc/raw_error.rs | 177 ++++++++++ .../src/jsonrpc/transport_actor.rs | 2 +- src/agent-client-protocol/src/lib.rs | 5 +- .../src/mcp_server/active_session.rs | 20 +- src/agent-client-protocol/src/role/acp.rs | 8 +- .../tests/jsonrpc_transport_close.rs | 6 +- .../tests/mcp_connector_errors.rs | 313 ++++++++++++++++++ .../tests/protocol_v2.rs | 26 +- .../tests/proxy_protocol_router_v2.rs | 10 +- .../tests/session_ordering.rs | 8 +- 22 files changed, 849 insertions(+), 138 deletions(-) create mode 100644 src/agent-client-protocol/src/jsonrpc/raw_error.rs create mode 100644 src/agent-client-protocol/tests/mcp_connector_errors.rs diff --git a/md/http-transport.md b/md/http-transport.md index 8f97b804..b613875d 100644 --- a/md/http-transport.md +++ b/md/http-transport.md @@ -67,6 +67,11 @@ can then queue until the GET attaches; an unknown-session GET returns 409 and does not allocate a mailbox. A batch registers all its session mailboxes before returning 202. +The server reserves inbound queue capacity before publishing that metadata. +A rejected POST therefore cannot roll back a mailbox another accepted POST +has adopted. Registration failure or cancellation before publication releases +the reserved queue slot without leaving partial metadata. + Keep the connection stream and existing session streams running during this setup so callback responses and cancellation can still progress. `HttpClient` does this automatically. When a `session/new` or `session/fork` response returns @@ -80,6 +85,11 @@ for the next frame. A WebSocket frame rejected by admission terminates that connection rather than silently losing input or waiting forever to drain a still-live agent. +A fatal outbound routing failure, such as mailbox overflow, closes the whole +connection and reclaims its agent and metadata even if the agent emits no more +messages. This is distinct from an ordinary SSE disconnect, which can reconnect +to the existing mailbox. + ## Features The crate does not enable either transport side by default. Opt into only the side(s) you need. diff --git a/md/migration-stateless-mcp.md b/md/migration-stateless-mcp.md index 39f03d31..844e83ac 100644 --- a/md/migration-stateless-mcp.md +++ b/md/migration-stateless-mcp.md @@ -117,6 +117,19 @@ or `try_send(frame)` for explicit fail-fast admission. The old `unbounded_send` API is removed; ignoring capacity errors silently loses protocol traffic. Finite queue and byte policies are configured through `ConnectionLimits`. +`RawJsonRpcMessage::Response` now carries the SDK's `RawJsonRpcResponse`, not +`schema::v1::Response`. Update raw response patterns to import +`agent_client_protocol::RawJsonRpcResponse`. Its `RawJsonRpcError` has an `i32` +code, `MaybeUndefined` data, and an extension map. Forward raw responses +unchanged to retain explicit null and unknown error fields. The error branch +stores `Box` so extensible errors do not enlarge every frame. + +`RawJsonRpcMessage::response(id, Result)` remains the convenience +constructor for ACP results. Other protocols should construct +`RawJsonRpcResponse` directly. Use `into_acp_error()` only when intentionally +dispatching an ACP error, not when forwarding MCP errors. Typed ACP request +consumers still receive the existing `Error` type. + ## Release checklist - Replace the draft Git schema pin with the released matching schema version. diff --git a/md/transport-architecture.md b/md/transport-architecture.md index 9f7764dc..fc5a1879 100644 --- a/md/transport-architecture.md +++ b/md/transport-architecture.md @@ -45,10 +45,17 @@ by the JSON-RPC envelope types from `agent-client-protocol-schema`: enum RawJsonRpcMessage { Request(Request), Notification(Notification), - Response(Response), + Response(RawJsonRpcResponse), } ``` +`RawJsonRpcResponse` uses the shared JSON-RPC response envelope with an opaque +JSON result and `RawJsonRpcError`. Raw errors keep numeric codes uninterpreted, +preserve unknown error fields, and distinguish omitted `data` from explicit +null. Only the typed ACP dispatcher converts them to ACP `Error`. Raw relays +and MCP connectors must not make that conversion: it would discard extensions +and apply the wrong protocol's error-code meaning. + At that boundary: - **Above**: Protocol layer works with application types (`OutgoingMessage`, `UntypedMessage`) diff --git a/src/agent-client-protocol-conductor/src/trace.rs b/src/agent-client-protocol-conductor/src/trace.rs index e83fa469..b3d6117f 100644 --- a/src/agent-client-protocol-conductor/src/trace.rs +++ b/src/agent-client-protocol-conductor/src/trace.rs @@ -12,10 +12,11 @@ use std::time::Instant; use agent_client_protocol::schema::SuccessorMessage; use agent_client_protocol::schema::v1::{ MessageMcpNotification, MessageMcpRequest, MessageMcpResponse, Notification as RpcNotification, - Request as RpcRequest, RequestId, Response as RpcResponse, + Request as RpcRequest, RequestId, }; use agent_client_protocol::{ - DynConnectTo, JsonRpcMessage, RawJsonRpcMessage, RawJsonRpcParams, Role, UntypedMessage, + DynConnectTo, JsonRpcMessage, RawJsonRpcMessage, RawJsonRpcParams, + RawJsonRpcResponse as RpcResponse, Role, UntypedMessage, }; use rustc_hash::FxHashMap; use serde::{Deserialize, Serialize}; diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index 791172f0..f43ccc3c 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -5,9 +5,8 @@ use std::{ use agent_client_protocol::{ Agent, BudgetedFrame, Channel, Client, ConnectTo, Error as AcpError, FrameAdmission, - FramePermit, FrameReceiver, FrameSender, RawJsonRpcMessage, TransportBatchEntry, - TransportFrame, - schema::v1::{RequestId, Response as RpcResponse}, + FramePermit, FrameReceiver, FrameSender, RawJsonRpcMessage, RawJsonRpcResponse as RpcResponse, + TransportBatchEntry, TransportFrame, schema::v1::RequestId, }; use async_tungstenite::tungstenite::Message as WsMessage; use futures::{ diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index f38504c9..07f0c0ba 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -5,8 +5,8 @@ use std::{ use agent_client_protocol::{ BudgetedFrame, Channel, ConnectionLimits, FrameAdmission, FramePermit, RawJsonRpcMessage, - TransportBatch, TransportBatchEntry, TransportFrame, - schema::v1::{RequestId, Response as RpcResponse}, + RawJsonRpcResponse as RpcResponse, TransportBatch, TransportBatchEntry, TransportFrame, + schema::v1::RequestId, }; use futures::{SinkExt, StreamExt}; use tokio::sync::{Mutex, RwLock, mpsc, watch}; @@ -169,26 +169,23 @@ impl Connection { .map_err(|_| "agent channel full or closed") } + pub(crate) fn reserve_inbound(&self) -> Result, &'static str> { + self.inbound_tx + .clone() + .try_reserve_owned() + .map_err(|_| "agent channel full or closed") + } + pub(crate) async fn register_post_routes( &self, sessions: &[String], routes: &[(RequestId, ResponseRoute)], permit: &FramePermit, - ) -> Result, &'static str> { + ) -> Result<(), &'static str> { if let OutboundTransport::Http(http) = &self.outbound_transport { http.register_post_routes(sessions, routes, permit).await } else { - Ok(Vec::new()) - } - } - - pub(crate) async fn rollback_post_routes( - &self, - sessions: &[String], - routes: &[(RequestId, ResponseRoute)], - ) { - if let OutboundTransport::Http(http) = &self.outbound_transport { - http.rollback_post_routes(sessions, routes).await; + Ok(()) } } @@ -403,7 +400,7 @@ impl HttpOutbound { sessions: &[String], routes: &[(RequestId, ResponseRoute)], permit: &FramePermit, - ) -> Result, &'static str> { + ) -> Result<(), &'static str> { // Lock both metadata tables in one order and check the whole batch // before inserting anything: rejection must never leave half a batch. let mut streams = self.session_streams.write().await; @@ -431,11 +428,8 @@ impl HttpOutbound { .map_err(|_| "HTTP session metadata capacity exceeded") }) .collect::, _>>()?; - for (id, session_permit) in new_sessions.iter().zip(session_permits) { - streams.insert( - id.clone(), - (Arc::new(OutboundMailbox::new()), Some(session_permit)), - ); + for (id, session_permit) in new_sessions.into_iter().zip(session_permits) { + streams.insert(id, (Arc::new(OutboundMailbox::new()), Some(session_permit))); } for (id, route) in routes { if let Some(id) = pending_route_key(id) { @@ -445,27 +439,7 @@ impl HttpOutbound { .push_back((route.clone(), Some(permit.clone()))); } } - Ok(new_sessions) - } - - async fn rollback_post_routes( - &self, - new_sessions: &[String], - routes: &[(RequestId, ResponseRoute)], - ) { - let mut streams = self.session_streams.write().await; - let mut pending = self.pending_routes.lock().await; - for id in new_sessions { - streams.remove(id); - } - for (id, _) in routes.iter().rev() { - if let Some(queue) = pending.get_mut(id) { - queue.pop_back(); - if queue.is_empty() { - pending.remove(id); - } - } - } + Ok(()) } #[cfg(test)] @@ -738,11 +712,36 @@ impl ConnectionRegistry { let (inbound_abort, inbound_abort_registration) = futures::future::AbortHandle::new_pair(); let inbound = futures::future::Abortable::new(inbound, inbound_abort_registration); let inbound_abort_for_outbound = inbound_abort.clone(); + let mut router_closed = closed_tx.subscribe(); let outbound = async move { - while let Some(msg) = agent_rx.next().await { - if outbound_tx.send(msg).await.is_err() { - inbound_abort_for_outbound.abort(); - break; + loop { + tokio::select! { + // A fatal router failure must tear down even an idle agent: + // waiting for another frame here can otherwise retain the + // agent and its registry entry forever. + changed = router_closed.changed() => { + if changed.is_err() || *router_closed.borrow() { + inbound_abort_for_outbound.abort(); + break; + } + } + msg = agent_rx.next() => { + let Some(msg) = msg else { break }; + let sent = tokio::select! { + sent = outbound_tx.send(msg) => sent.is_ok(), + changed = router_closed.changed() => { + if changed.is_err() || *router_closed.borrow() { + false + } else { + continue; + } + } + }; + if !sent { + inbound_abort_for_outbound.abort(); + break; + } + } } } }; @@ -1253,12 +1252,10 @@ mod tests { assert!(matches!( frame.frame(), - TransportFrame::Single(RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { - id: RequestId::Number(1), - .. - } - )) + TransportFrame::Single(RawJsonRpcMessage::Response(RpcResponse::Result { + id: RequestId::Number(1), + .. + })) )); timeout(Duration::from_secs(1), async { loop { @@ -1306,6 +1303,87 @@ mod tests { )); } + #[tokio::test] + async fn fatal_router_overflow_tears_down_idle_agent_and_metadata() { + struct BurstThenIdle(Arc); + struct Dropped(Arc); + impl Drop for Dropped { + fn drop(&mut self) { + self.0.store(true, std::sync::atomic::Ordering::SeqCst); + } + } + impl AgentFactory for BurstThenIdle { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (agent, transport) = Channel::duplex(); + let dropped = Dropped(self.0.clone()); + let future = Box::pin(async move { + let _dropped = dropped; + for _ in 0..33 { + agent + .tx + .send_frame(TransportFrame::Single( + RawJsonRpcMessage::notification( + "test/burst".into(), + serde_json::json!({}), + ) + .unwrap(), + )) + .await + .unwrap(); + } + std::future::pending::<()>().await; + Ok(()) + }); + (transport, future) + } + } + + let dropped = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let registry = ConnectionRegistry::new(Arc::new(BurstThenIdle(dropped.clone()))); + let (id, connection) = registry.create_connection().await; + let source = connection + .admit_frame_to_agent(TransportFrame::Single( + RawJsonRpcMessage::notification("test/source".into(), serde_json::json!({})) + .unwrap(), + )) + .unwrap(); + connection + .register_post_routes( + &["S".into()], + &[(RequestId::Number(1), ResponseRoute::Session("S".into()))], + source.permit(), + ) + .await + .unwrap(); + let OutboundTransport::Http(http) = &connection.outbound_transport else { + unreachable!() + }; + assert_eq!(http.pending_routes.lock().await.len(), 1); + drop(source); + connection.start_router().await; + timeout(Duration::from_secs(1), async { + let mut closed = connection.subscribe_closed(); + while !*closed.borrow() { + closed.changed().await.unwrap(); + } + while registry.get(&id).await.is_some() + || !dropped.load(std::sync::atomic::Ordering::SeqCst) + || !http.pending_routes.lock().await.is_empty() + || connection.subscribe_session_stream("S").await.is_some() + { + sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("fatal router exit must stop idle agent and remove registry entry"); + assert!(http.pending_routes.lock().await.is_empty()); + } + #[tokio::test] async fn protocol_level_notification_routes_to_connection_stream() { let exit = Arc::new(Notify::new()); diff --git a/src/agent-client-protocol-http/src/connection_admission_tests.rs b/src/agent-client-protocol-http/src/connection_admission_tests.rs index dfe70f65..c0f6838f 100644 --- a/src/agent-client-protocol-http/src/connection_admission_tests.rs +++ b/src/agent-client-protocol-http/src/connection_admission_tests.rs @@ -154,26 +154,55 @@ async fn route_and_session_metadata_admission_is_atomic_and_releases_permits() { } #[tokio::test] -async fn rolling_back_rejected_transport_send_removes_only_new_metadata() { +async fn concurrent_post_reservations_preserve_adopted_session() { let frame = TransportFrame::Single( RawJsonRpcMessage::request("test/request".into(), json!({}), RequestId::Number(1)).unwrap(), ); let (_caller, transport) = Channel::duplex(); - let (_, permit) = transport - .tx - .admission() - .try_admit(frame) - .unwrap() - .into_parts(); - let http = HttpOutbound::new(); - let routes = [(RequestId::Number(1), ResponseRoute::Session("one".into()))]; - let new_sessions = http - .register_post_routes(&["one".into()], &routes, &permit) + let (inbound_tx, mut inbound_rx) = mpsc::channel(2); + let connection = Connection { + inbound_tx, + inbound_admission: transport.tx.admission(), + outbound_rx: Mutex::new(None), + agent_handle: Mutex::new(None), + router_handle: Mutex::new(None), + closed_tx: watch::channel(false).0, + outbound_transport: OutboundTransport::http(), + }; + let a = connection.reserve_inbound().unwrap(); + let b = connection.reserve_inbound().unwrap(); + assert!(connection.reserve_inbound().is_err()); + + // A publishes S first; B adopts S and enqueues while A is paused. + // Neither commit may subsequently fail queue admission or erase S. + let a_frame = connection.admit_frame_to_agent(frame.clone()).unwrap(); + connection + .register_post_routes(&["S".into()], &[], a_frame.permit()) .await .unwrap(); - http.rollback_post_routes(&new_sessions, &routes).await; - assert!(http.pending_routes.lock().await.is_empty()); - assert!(http.session_streams.read().await.is_empty()); + let OutboundTransport::Http(http) = &connection.outbound_transport else { + unreachable!("HTTP test connection"); + }; + let original = http.session_streams.read().await["S"].0.clone(); + let b_frame = connection.admit_frame_to_agent(frame).unwrap(); + connection + .register_post_routes(&["S".into()], &[], b_frame.permit()) + .await + .unwrap(); + assert!(Arc::ptr_eq( + &original, + &http.session_streams.read().await["S"].0 + )); + b.send(b_frame); + a.send(a_frame); + assert!(inbound_rx.recv().await.is_some()); + assert!(inbound_rx.recv().await.is_some()); + assert!(connection.subscribe_session_stream("S").await.is_some()); + + // A cancelled before metadata registration cannot strand a queue slot. + let cancelled = connection.reserve_inbound().unwrap(); + drop(cancelled); + assert!(connection.reserve_inbound().is_ok()); } #[tokio::test] diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index a3d6676d..bda0274a 100644 --- a/src/agent-client-protocol-http/src/http_server.rs +++ b/src/agent-client-protocol-http/src/http_server.rs @@ -1,8 +1,8 @@ use std::{convert::Infallible, error::Error as _, sync::Arc, time::Duration}; use agent_client_protocol::{ - RawJsonRpcMessage, TransportBatchEntry, TransportFrame, schema::v1::RequestId, - schema::v1::Response as RpcResponse, + RawJsonRpcMessage, RawJsonRpcResponse as RpcResponse, TransportBatchEntry, TransportFrame, + schema::v1::RequestId, }; use axum::{ body::Body, @@ -179,21 +179,22 @@ pub(crate) async fn handle_post( Ok(frame) => frame, Err(error) => return (StatusCode::TOO_MANY_REQUESTS, error).into_response(), }; + // Claim the queue slot before publishing session and response-route + // metadata. After publication, sending through this permit cannot fail + // due to another POST filling the queue. + let inbound_slot = match connection.reserve_inbound() { + Ok(slot) => slot, + Err(error) => return (StatusCode::TOO_MANY_REQUESTS, error).into_response(), + }; let permit = admitted.permit().clone(); - let new_sessions = match connection + if let Err(error) = connection .register_post_routes(&session_routes, &pending_routes, &permit) .await { - Ok(new_sessions) => new_sessions, - Err(error) => return (StatusCode::TOO_MANY_REQUESTS, error).into_response(), - }; - drop(permit); - if connection.send_budgeted_frame_to_agent(admitted).is_err() { - connection - .rollback_post_routes(&new_sessions, &pending_routes) - .await; - return StatusCode::INTERNAL_SERVER_ERROR.into_response(); + return (StatusCode::TOO_MANY_REQUESTS, error).into_response(); } + drop(permit); + inbound_slot.send(admitted); connection.cancel_pending_routes(&cancellations).await; StatusCode::ACCEPTED.into_response() } @@ -541,6 +542,54 @@ mod tests { } } + #[tokio::test] + async fn rejected_post_does_not_remove_accepted_session_stream() { + let (forwarded, _receiver) = mpsc::unbounded_channel(); + let registry = Arc::new(ConnectionRegistry::new(Arc::new(CapturingAgentFactory { + forwarded, + }))); + let (id, connection) = registry.create_connection().await; + let post = || { + Request::builder() + .method("POST") + .uri("/acp") + .header(header::CONTENT_TYPE, JSON_MIME_TYPE) + .header(HEADER_CONNECTION_ID, id.as_str()) + .header(HEADER_SESSION_ID, "S") + .body(Body::from( + json!({"jsonrpc":"2.0","method":"session/update","params":{}}).to_string(), + )) + .unwrap() + }; + assert_eq!( + handle_post(State(registry.clone()), post()).await.status(), + StatusCode::ACCEPTED + ); + let mut slots = Vec::new(); + while let Ok(slot) = connection.reserve_inbound() { + slots.push(slot); + } + assert!(!slots.is_empty()); + assert_eq!( + handle_post(State(registry.clone()), post()).await.status(), + StatusCode::TOO_MANY_REQUESTS + ); + let get = Request::builder() + .uri("/acp") + .header(header::ACCEPT, EVENT_STREAM_MIME_TYPE) + .header(HEADER_CONNECTION_ID, id.as_str()) + .header(HEADER_SESSION_ID, "S") + .body(Body::empty()) + .unwrap(); + assert_eq!( + handle_get(registry.clone(), get).await.status(), + StatusCode::OK + ); + drop(slots); + registry.remove(&id).await; + connection.shutdown().await; + } + struct RejectingInitializeAgentFactory; impl AgentFactory for RejectingInitializeAgentFactory { diff --git a/src/agent-client-protocol-http/src/websocket_server.rs b/src/agent-client-protocol-http/src/websocket_server.rs index 4e614aab..08356125 100644 --- a/src/agent-client-protocol-http/src/websocket_server.rs +++ b/src/agent-client-protocol-http/src/websocket_server.rs @@ -245,8 +245,8 @@ where #[cfg(test)] mod tests { use agent_client_protocol::{ - BudgetedFrame, Channel, TransportBatch, TransportBatchEntry, TransportFrame, - schema::v1::{RequestId, Response as RpcResponse}, + BudgetedFrame, Channel, RawJsonRpcResponse as RpcResponse, TransportBatch, + TransportBatchEntry, TransportFrame, schema::v1::RequestId, }; use async_tungstenite::{tokio::connect_async, tungstenite::Message as ClientWsMessage}; use axum::{Router, extract::WebSocketUpgrade, routing::get}; diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index 35f65fac..2ad091a2 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -7,12 +7,20 @@ - Target MCP 2026-07-28 with server-addressed `mcp/message` operations and logical `McpRequestId`s. Remove connect/disconnect and reverse MCP requests; providers send request-scoped notifications and use ACP cancellation. -- Create an independent backend per operation and expose `request_id()` in - attached MCP contexts instead of `connection_id()`. Preserve standalone MCP - serving independently of the unstable ACP transport feature. +- Use reusable native services with owned per-operation execution and cleanup; + retain an explicit connector-backed adapter for factory-based servers. + Expose `request_id()` in attached MCP contexts instead of `connection_id()`. + Preserve standalone serving independently of the unstable ACP feature. - Validate modern request metadata, restrict discovery to the binding's MCP - revision, and add native admission/payload limits. End-to-end native queue - backpressure remains required before stabilization. + revision, and bound native admission, payloads, and transport queues. +- Preserve MCP error codes, omitted/null data, and extensions in a distinct + inner outcome carrier, including connector-backed byte-stream servers. + +### Changed (breaking transport APIs) + +- Raw JSON-RPC responses use `RawJsonRpcResponse` and `RawJsonRpcError`, + preserving error fields without ACP interpretation. Typed ACP consumers + continue receiving `Error`; raw adapters must use the new response type. ### Added diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index f71bbf63..280c4da1 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -2,7 +2,7 @@ use agent_client_protocol_schema::v1::{ JsonRpcMessage as VersionedJsonRpcMessage, Notification as RpcNotification, - Request as RpcRequest, RequestId, Response as RpcResponse, SessionId, + Request as RpcRequest, RequestId, SessionId, }; // Types re-exported from crate root @@ -33,6 +33,7 @@ pub(crate) mod handlers; mod incoming_actor; mod outgoing_actor; mod protocol_compat; +mod raw_error; pub(crate) mod run; mod task_actor; mod transport_actor; @@ -45,6 +46,7 @@ use crate::jsonrpc::handlers::{ChainedHandler, NamedHandler}; use crate::jsonrpc::handlers::{MessageHandler, NotificationHandler, RequestHandler}; use crate::jsonrpc::outgoing_actor::{OutgoingMessageTx, send_raw_message}; use crate::jsonrpc::protocol_compat::{ProtocolCompat, ProtocolMode}; +pub use crate::jsonrpc::raw_error::{RawJsonRpcError, RawJsonRpcResponse}; use crate::jsonrpc::run::SpawnedRun; use crate::jsonrpc::run::{ChainRun, NullRun, RunWithConnectionTo}; use crate::jsonrpc::task_actor::{Task, TaskTx}; @@ -64,8 +66,8 @@ pub enum RawJsonRpcMessage { Request(RpcRequest), /// A JSON-RPC notification without a response. Notification(RpcNotification), - /// A JSON-RPC response to a prior request. - Response(RpcResponse), + /// A response with an opaque result or transport-neutral error object. + Response(RawJsonRpcResponse), } /// A JSON-RPC frame exchanged between protocol components and transports. @@ -408,19 +410,25 @@ impl RawJsonRpcMessage { })) } - /// Build a raw JSON-RPC response message. + /// Build a raw response from an ACP result. + /// + /// For other protocols, construct [`RawJsonRpcResponse`] directly so error + /// codes and fields are not first interpreted as ACP errors. #[must_use] pub fn response(id: RequestId, response: Result) -> Self { - Self::Response(RpcResponse::new(id, response)) + Self::Response(RawJsonRpcResponse::new( + id, + response.map_err(|error| Box::new(error.into())), + )) } /// The response id, if this is a response. #[must_use] pub fn response_id(&self) -> Option<&RequestId> { match self { - Self::Response(RpcResponse::Result { id, .. } | RpcResponse::Error { id, .. }) => { - Some(id) - } + Self::Response( + RawJsonRpcResponse::Result { id, .. } | RawJsonRpcResponse::Error { id, .. }, + ) => Some(id), Self::Request(_) | Self::Notification(_) => None, } } @@ -477,11 +485,10 @@ impl<'de> Deserialize<'de> for RawJsonRpcMessage { Ok(Self::Notification(notification)) } } else if !has_method && has_id && has_result != has_error { - let response = serde_json::from_value::< - VersionedJsonRpcMessage>, - >(value) - .map_err(serde::de::Error::custom)? - .into_inner(); + let response = + serde_json::from_value::>(value) + .map_err(serde::de::Error::custom)? + .into_inner(); Ok(Self::Response(response)) } else { Err(serde::de::Error::custom("invalid JSON-RPC message")) diff --git a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs index a8ff2df3..0fe5cdb6 100644 --- a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs @@ -17,6 +17,7 @@ use crate::jsonrpc::PendingReplies; use crate::jsonrpc::PendingReply; use crate::jsonrpc::RawJsonRpcMessage; use crate::jsonrpc::RawJsonRpcParams; +use crate::jsonrpc::RawJsonRpcResponse as Response; use crate::jsonrpc::RequestReplyTarget; use crate::jsonrpc::Responder; use crate::jsonrpc::ResponseDestination; @@ -31,7 +32,7 @@ use crate::jsonrpc::{BudgetedFrame, FramePermit, TransportFrame}; use crate::jsonrpc::{is_response_only_shape, raw_is_response_only_shape}; use crate::role::Role; -use crate::schema::v1::{RequestId, Response}; +use crate::schema::v1::RequestId; use super::Handled; @@ -210,7 +211,7 @@ pub(super) async fn incoming_protocol_actor( Ok(RawJsonRpcMessage::Response(response)) => { let (id, result) = match response { Response::Result { id, result } => (id, Ok(result)), - Response::Error { id, error } => (id, Err(error)), + Response::Error { id, error } => (id, Err(error.into_acp_error())), }; tracing::trace!(?id, "Handling response"); diff --git a/src/agent-client-protocol/src/jsonrpc/raw_error.rs b/src/agent-client-protocol/src/jsonrpc/raw_error.rs new file mode 100644 index 00000000..00a6ec97 --- /dev/null +++ b/src/agent-client-protocol/src/jsonrpc/raw_error.rs @@ -0,0 +1,177 @@ +//! Transport-level error objects, before choosing an application protocol. + +use agent_client_protocol_schema::MaybeUndefined; +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; + +/// A JSON-RPC error without ACP-specific interpretation. +/// +/// Raw transports and relays preserve unknown fields and distinguish omitted +/// `data` from explicit JSON null. Convert to [`crate::Error`] only when +/// dispatching an ACP response; an MCP error code belongs to a different domain. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[non_exhaustive] +pub struct RawJsonRpcError { + /// The peer's numeric error code, not an ACP [`crate::ErrorCode`]. + pub code: i32, + /// The peer's error message. + pub message: String, + /// Optional error data. Explicit null is retained separately from omission. + #[serde(default, skip_serializing_if = "MaybeUndefined::is_undefined")] + pub data: MaybeUndefined, + /// Additional fields on the error object. + #[serde(flatten)] + pub extra: Map, +} + +/// A transport-level JSON-RPC response with an opaque result or raw error. +/// +/// Errors are boxed so their extensible representation does not enlarge every +/// request, notification, and queued frame. +pub type RawJsonRpcResponse = + agent_client_protocol_schema::rpc::Response>; + +impl RawJsonRpcError { + /// Construct an error without data or extension fields. + #[must_use] + pub fn new(code: i32, message: impl Into) -> Self { + Self { + code, + message: message.into(), + data: MaybeUndefined::Undefined, + extra: Map::new(), + } + } + + /// Set error data, preserving explicit null. + #[must_use] + pub fn data(mut self, data: Value) -> Self { + self.data = if data.is_null() { + MaybeUndefined::Null + } else { + MaybeUndefined::Value(data) + }; + self + } + + /// Interpret this error as an ACP response for the typed dispatcher. + /// + /// ACP's error type does not model extension fields, so this intentionally + /// discards `extra`. Do not use it when forwarding raw frames or projecting + /// errors from another protocol such as MCP. + #[must_use] + pub fn into_acp_error(self) -> crate::Error { + let mut error = crate::Error::new(self.code, self.message); + error.data = match self.data { + MaybeUndefined::Undefined => None, + MaybeUndefined::Null => Some(Value::Null), + MaybeUndefined::Value(data) => Some(data), + }; + error + } +} + +impl From for RawJsonRpcError { + fn from(error: crate::Error) -> Self { + let raw = Self::new(error.code.into(), error.message); + match error.data { + Some(data) => raw.data(data), + None => raw, + } + } +} + +impl std::fmt::Display for RawJsonRpcError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(formatter, "{} ({})", self.message, self.code) + } +} + +impl std::error::Error for RawJsonRpcError {} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{Channel, RawJsonRpcMessage, TransportFrame}; + use futures::{SinkExt as _, StreamExt as _}; + use serde_json::json; + + #[test] + fn raw_errors_preserve_omission_null_values_and_extensions() { + for data in [None, Some(Value::Null), Some(json!({"detail":[1,2]}))] { + let mut error = json!({ + "code": -32000, + "message": "peer", + "extension": {"retry": true}, + "_meta": {"opaque": "kept"} + }); + if let Some(data) = &data { + error["data"] = data.clone(); + } + let wire = json!({"jsonrpc":"2.0", "id":"logical", "error":error}); + let parsed: RawJsonRpcMessage = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(serde_json::to_value(&parsed).unwrap(), wire); + let RawJsonRpcMessage::Response(RawJsonRpcResponse::Error { error, .. }) = parsed + else { + panic!("expected a raw error response"); + }; + assert_eq!(error.code, -32000); + match data { + None => assert!(error.data.is_undefined()), + Some(Value::Null) => assert!(error.data.is_null()), + Some(value) => assert_eq!(error.data.value(), Some(&value)), + } + } + } + + #[tokio::test] + async fn raw_error_batch_survives_framing_and_budgeted_relay() { + let wire = json!([ + {"jsonrpc":"2.0", "id":"error", "error":{ + "code":-32000, "message":"peer", "data":null, "extension":{"retry":true} + }}, + {"jsonrpc":"2.0", "id":"success", "result":null} + ]); + let frame = TransportFrame::parse_json(&wire.to_string()); + assert!(matches!(&frame, TransportFrame::Batch(_))); + let (source, mut relay_in) = Channel::duplex(); + let (mut relay_out, mut destination) = Channel::duplex(); + source.tx.send_frame(frame).await.unwrap(); + let admitted = relay_in.rx.next().await.unwrap(); + relay_out.tx.send(admitted).await.unwrap(); + let received = destination.rx.next().await.unwrap(); + let received: Value = serde_json::from_str(&received.frame().to_json().unwrap()).unwrap(); + assert_eq!(received, wire); + } + + #[test] + fn acp_error_interpretation_is_explicit_and_keeps_data_presence() { + let raw = RawJsonRpcError::new(-32000, "peer"); + assert_eq!(raw.clone().into_acp_error().data, None); + let error = raw.data(Value::Null).into_acp_error(); + assert_eq!(error.code, crate::ErrorCode::AuthRequired); + assert_eq!(error.data, Some(Value::Null)); + let roundtrip = RawJsonRpcError::from(error); + assert!(roundtrip.data.is_null()); + assert!(roundtrip.extra.is_empty()); + } + + #[test] + fn malformed_raw_errors_are_still_rejected() { + for error in [ + Value::Null, + json!({"code":-32000}), + json!({"message":"peer"}), + json!({"code":null, "message":"peer"}), + json!({"code":1.5, "message":"peer"}), + json!({"code":-32000, "message":null}), + ] { + assert!( + serde_json::from_value::( + json!({"jsonrpc":"2.0", "id":1, "error":error}) + ) + .is_err() + ); + } + } +} diff --git a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs index be2cb522..d71d4004 100644 --- a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs @@ -1,8 +1,8 @@ use std::pin::pin; // Types re-exported from crate root +use crate::RawJsonRpcResponse as Response; use crate::jsonrpc::{RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame}; -use crate::schema::v1::Response; use futures::StreamExt as _; use serde::Deserialize as _; diff --git a/src/agent-client-protocol/src/lib.rs b/src/agent-client-protocol/src/lib.rs index 6d280fc2..7f58393e 100644 --- a/src/agent-client-protocol/src/lib.rs +++ b/src/agent-client-protocol/src/lib.rs @@ -148,8 +148,9 @@ pub use jsonrpc::{ FrameSender, HandleConnectionClose, HandleDispatchFrom, Handled, INCOMING_TRANSPORT_CLOSED_REASON, IntoHandled, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, Lines, NullClose, NullHandler, RawConnectionContext, - RawJsonRpcMessage, RawJsonRpcParams, Responder, ResponseRouter, SentRequest, TransportBatch, - TransportBatchEntry, TransportFrame, UntypedMessage, is_incoming_transport_closed, + RawJsonRpcError, RawJsonRpcMessage, RawJsonRpcParams, RawJsonRpcResponse, Responder, + ResponseRouter, SentRequest, TransportBatch, TransportBatchEntry, TransportFrame, + UntypedMessage, is_incoming_transport_closed, run::{ChainRun, NullRun, RunWithConnectionTo}, }; pub use jsonrpc::{RequestCancellation, is_cancel_request_notification}; diff --git a/src/agent-client-protocol/src/mcp_server/active_session.rs b/src/agent-client-protocol/src/mcp_server/active_session.rs index d0ed1b4f..e910434a 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -15,8 +15,8 @@ use std::{ use crate::{ Agent, Channel, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, RawJsonRpcMessage, Responder, Role, - TransportFrame, + JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, RawJsonRpcError, RawJsonRpcMessage, + RawJsonRpcResponse, Responder, Role, TransportFrame, mcp_server::{ MCP_BACKEND_FAILURE, MCP_RESOURCE_EXHAUSTED, McpConnectionContext, McpConnectionTo, McpOperationCancellation, McpOutcome, McpRequest, McpRequestContext, McpServerConnect, @@ -88,11 +88,11 @@ impl McpProtocol for V1McpProtocol { } } -fn into_mcp_error(error: crate::Error) -> McpError { - let mut mcp = McpError::new(error.code.into(), error.message); - if let Some(data) = error.data { - mcp = mcp.data(data); - } +fn into_mcp_error(error: impl Into) -> McpError { + let error = error.into(); + let mut mcp = McpError::new(error.code, error.message); + mcp.data = error.data; + mcp.extra = error.extra; mcp } @@ -482,11 +482,11 @@ where RawJsonRpcMessage::Response(response) => { // Returning ends notification forwarding before the terminal reply. return match response { - crate::schema::v1::Response::Result { result, .. } => { + RawJsonRpcResponse::Result { result, .. } => { Ok(McpOutcome::Result(result)) } - crate::schema::v1::Response::Error { error, .. } => { - Ok(McpOutcome::Error(into_mcp_error(error))) + RawJsonRpcResponse::Error { error, .. } => { + Ok(McpOutcome::Error(into_mcp_error(*error))) } }; } diff --git a/src/agent-client-protocol/src/role/acp.rs b/src/agent-client-protocol/src/role/acp.rs index ba03d6e4..067bc045 100644 --- a/src/agent-client-protocol/src/role/acp.rs +++ b/src/agent-client-protocol/src/role/acp.rs @@ -17,16 +17,16 @@ use crate::role::{HasPeer, RemoteStyle}; #[cfg(not(feature = "unstable_protocol_v2"))] use crate::schema::InitializeProxyRequest; use crate::schema::METHOD_INITIALIZE_PROXY; +#[cfg(feature = "unstable_protocol_v2")] +use crate::schema::v1::RequestId; use crate::schema::v1::{InitializeRequest, SessionId}; #[cfg(not(feature = "unstable_protocol_v2"))] use crate::schema::v1::{NewSessionRequest, NewSessionResponse}; #[cfg(feature = "unstable_protocol_v2")] -use crate::schema::v1::{RequestId, Response as RpcResponse}; -#[cfg(feature = "unstable_protocol_v2")] use crate::schema::{ProtocolVersion, v2}; use crate::util::MatchDispatchFrom; #[cfg(feature = "unstable_protocol_v2")] -use crate::{Channel, RawJsonRpcMessage, RawJsonRpcParams}; +use crate::{Channel, RawJsonRpcMessage, RawJsonRpcParams, RawJsonRpcResponse as RpcResponse}; use crate::{ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, Role, RoleId}; #[cfg(feature = "unstable_protocol_v2")] @@ -1134,7 +1134,7 @@ impl InitializeResponse { }), RawJsonRpcMessage::Response(RpcResponse::Error { id, error }) => Ok(Self { id, - result: Err(error), + result: Err(error.into_acp_error()), }), message => Err(crate::Error::invalid_request().data(format!( "first ACP response must be an initialize response, got {message:?}", diff --git a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs index b19e929d..7e80ec64 100644 --- a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs +++ b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs @@ -13,10 +13,10 @@ use std::{ use agent_client_protocol::{ BudgetedFrame, ByteStreams, Channel, ConnectTo, ConnectionTo, Dispatch, Error, Handled, - JsonRpcMessage, JsonRpcRequest, Lines, RawJsonRpcMessage, TransportFrame, UntypedMessage, - is_incoming_transport_closed, + JsonRpcMessage, JsonRpcRequest, Lines, RawJsonRpcMessage, RawJsonRpcResponse as Response, + TransportFrame, UntypedMessage, is_incoming_transport_closed, role::{Role, UntypedRole}, - schema::v1::{RequestId, Response}, + schema::v1::RequestId, }; use agent_client_protocol_test::{MyRequest, MyResponse}; use futures::{FutureExt as _, StreamExt as _, future::join, stream}; diff --git a/src/agent-client-protocol/tests/mcp_connector_errors.rs b/src/agent-client-protocol/tests/mcp_connector_errors.rs new file mode 100644 index 00000000..6865f44b --- /dev/null +++ b/src/agent-client-protocol/tests/mcp_connector_errors.rs @@ -0,0 +1,313 @@ +#![cfg(feature = "unstable_mcp_over_acp")] + +use std::time::Duration; + +#[cfg(feature = "unstable_protocol_v2")] +use agent_client_protocol::V2ConnectionTo; +#[cfg(feature = "unstable_protocol_v2")] +use agent_client_protocol::schema::v2; +use agent_client_protocol::{ + Agent, ByteStreams, Client, ConnectTo, ConnectionTo, DynConnectTo, Error, Responder, + RunWithConnectionTo, + mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, + role, + schema::v1, +}; +use serde_json::{Map, Value, json}; +use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader, duplex}; +use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; + +const TIMEOUT: Duration = Duration::from_secs(10); +const CASES: &[(&str, &str)] = &[ + ( + "absent", + r#"{"code":-32000,"message":"backend failed","extension":{"retry":false}}"#, + ), + ( + "null", + r#"{"code":-32000,"message":"backend failed","data":null,"extension":{"retry":false}}"#, + ), + ( + "object", + r#"{"code":-32000,"message":"backend failed","data":{"cause":"upstream"},"extension":{"retry":false}}"#, + ), +]; + +struct WireConnector; + +impl McpServerConnect for WireConnector { + fn name(&self) -> String { + "raw-wire-errors".into() + } + + fn connect(&self, context: McpConnectionTo) -> DynConnectTo { + assert!( + context.request_id().is_some(), + "expected an ACP MCP request" + ); + DynConnectTo::new(WireBackend) + } +} + +struct WireBackend; + +impl ConnectTo for WireBackend { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + let (sdk_output, peer_input) = duplex(4096); + let (peer_output, sdk_input) = duplex(4096); + let transport = ByteStreams::new(sdk_output.compat_write(), sdk_input.compat()); + let peer = async move { + let mut reader = BufReader::new(peer_input); + let mut line = String::new(); + reader + .read_line(&mut line) + .await + .expect("read MCP wire request"); + let request: Value = serde_json::from_str(&line).expect("valid MCP wire request"); + assert_eq!(request["jsonrpc"], "2.0"); + assert_eq!( + request["params"]["_meta"]["io.modelcontextprotocol/protocolVersion"], + "2026-07-28" + ); + let id = serde_json::to_string(&request["id"]).unwrap(); + let method = request["method"].as_str().expect("MCP method"); + let response = if method == "success" { + assert_eq!( + request["id"], "next", + "backend must see the logical request ID" + ); + format!(r#"{{"jsonrpc":"2.0","id":{id},"result":null}}"#) + } else { + let index = CASES + .iter() + .position(|(case, _)| case == &method) + .expect("known test case"); + assert_eq!( + request["id"], + format!("wire-{index}"), + "backend must see the logical request ID" + ); + let (_, error) = CASES + .iter() + .find(|(case, _)| case == &method) + .expect("known test case"); + // Literal backend JSON, not ACP Error or RawJsonRpcMessage::response: + // those typed constructors cannot represent data:null or extension. + format!(r#"{{"jsonrpc":"2.0","id":{id},"error":{error}}}"#) + }; + let mut output = peer_output; + output + .write_all(response.as_bytes()) + .await + .expect("write MCP response"); + output.write_all(b"\n").await.expect("frame MCP response"); + output.shutdown().await.expect("close MCP response stream"); + }; + let (result, ()) = tokio::join!(client.connect_to(transport), peer); + result + } +} + +struct IdleRunner; + +impl RunWithConnectionTo for IdleRunner { + async fn run_with_connection_to(self, _connection: ConnectionTo) -> Result<(), Error> { + std::future::pending().await + } +} + +#[cfg(feature = "unstable_protocol_v2")] +fn cwd() -> std::path::PathBuf { + std::env::current_dir().expect("cwd") +} + +fn params() -> Map { + serde_json::from_value(json!({ + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } + })) + .unwrap() +} + +fn assert_error(error: impl serde::Serialize, case: &str) { + let value = serde_json::to_value(error).unwrap(); + assert_eq!(value["code"], -32000, "{case}"); + assert_eq!(value["message"], "backend failed", "{case}"); + assert_eq!(value["extension"], json!({"retry": false}), "{case}"); + let object = value.as_object().unwrap(); + match case { + "absent" => assert!(!object.contains_key("data"), "{value}"), + "null" => assert_eq!(object.get("data"), Some(&Value::Null)), + "object" => assert_eq!(object.get("data"), Some(&json!({"cause":"upstream"}))), + _ => unreachable!(), + } +} + +async fn v1_requests( + connection: ConnectionTo, + server_id: v1::McpServerAcpId, +) -> Result<(), Error> { + for (index, (case, _)) in CASES.iter().enumerate() { + let response = connection + .send_request( + v1::MessageMcpRequest::new(server_id.clone(), format!("wire-{index}"), *case) + .params(params()), + ) + .block_task() + .await?; + match response { + v1::MessageMcpResponse::Error { error, .. } => assert_error(error, case), + other => panic!("{case}: expected inner MCP error, got {other:?}"), + } + } + let response = connection + .send_request(v1::MessageMcpRequest::new(server_id, "next", "success").params(params())) + .block_task() + .await?; + assert!(matches!( + response, + v1::MessageMcpResponse::Result { + result: Value::Null, + .. + } + )); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn v1_connector_preserves_backend_wire_errors() { + let (done_tx, done_rx) = tokio::sync::oneshot::channel(); + let done_tx = std::sync::Mutex::new(Some(done_tx)); + let agent = Agent.builder().on_receive_request( + async move |request: v1::NewSessionRequest, + responder: Responder, + connection: ConnectionTo| { + let [v1::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected one native MCP server") + }; + let server_id = server.server_id.clone(); + responder.respond(v1::NewSessionResponse::new(v1::SessionId::new( + "wire-session", + )))?; + let done = done_tx.lock().unwrap().take().expect("one setup request"); + let requests = connection.clone(); + connection.spawn(async move { + drop(done.send(v1_requests(requests, server_id).await)); + Ok(()) + }) + }, + agent_client_protocol::on_receive_request!(), + ); + let test = Client.builder().connect_with(agent, async |connection| { + connection + .build_session_cwd()? + .with_mcp_server(McpServer::::new(WireConnector, IdleRunner))? + .block_task() + .run_until(async |_session| { + done_rx.await.map_err(Error::into_internal_error)??; + Ok(()) + }) + .await?; + Ok(()) + }); + tokio::time::timeout(TIMEOUT, test) + .await + .expect("v1 connector timed out") + .expect("v1 connector failed"); +} + +#[cfg(feature = "unstable_protocol_v2")] +async fn v2_requests( + connection: V2ConnectionTo, + server_id: v2::McpServerAcpId, +) -> Result<(), Error> { + for (index, (case, _)) in CASES.iter().enumerate() { + let response = connection + .send_request( + v2::MessageMcpRequest::new(server_id.clone(), format!("wire-{index}"), *case) + .params(params()), + ) + .block_task() + .await?; + match response { + v2::MessageMcpResponse::Error { error, .. } => assert_error(error, case), + other => panic!("{case}: expected inner MCP error, got {other:?}"), + } + } + let response = connection + .send_request(v2::MessageMcpRequest::new(server_id, "next", "success").params(params())) + .block_task() + .await?; + assert!(matches!( + response, + v2::MessageMcpResponse::Result { + result: Value::Null, + .. + } + )); + Ok(()) +} + +#[cfg(feature = "unstable_protocol_v2")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn v2_connector_preserves_backend_wire_errors() { + let (done_tx, done_rx) = tokio::sync::oneshot::channel(); + let done_tx = std::sync::Mutex::new(Some(done_tx)); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _connection: V2ConnectionTo| { + responder.respond(v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("wire-backend-test", env!("CARGO_PKG_VERSION")), + )) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + connection: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected one native MCP server") + }; + let server_id = server.server_id.clone(); + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new( + "wire-session", + )))?; + let done = done_tx.lock().unwrap().take().expect("one setup request"); + let requests = connection.clone(); + connection.spawn(async move { + drop(done.send(v2_requests(requests, server_id).await)); + Ok(()) + }) + }, + agent_client_protocol::on_receive_request!(), + ); + let test = Client.v2().connect_with(agent, async |connection| { + connection + .send_request(v2::InitializeRequest::new( + agent_client_protocol::schema::ProtocolVersion::V2, + v2::Implementation::new("wire-backend-test", env!("CARGO_PKG_VERSION")), + )) + .block_task() + .await?; + let session = connection + .build_session_from(v2::NewSessionRequest::new(cwd())) + .with_mcp_server(McpServer::::new(WireConnector, IdleRunner))? + .start_session() + .block_task() + .await?; + done_rx.await.map_err(Error::into_internal_error)??; + drop(session); + Ok(()) + }); + tokio::time::timeout(TIMEOUT, test) + .await + .expect("v2 connector timed out") + .expect("v2 connector failed"); +} diff --git a/src/agent-client-protocol/tests/protocol_v2.rs b/src/agent-client-protocol/tests/protocol_v2.rs index 562b2d33..8eeab227 100644 --- a/src/agent-client-protocol/tests/protocol_v2.rs +++ b/src/agent-client-protocol/tests/protocol_v2.rs @@ -421,7 +421,11 @@ impl ConnectTo for FutureInitializeV2Client { let TransportFrame::Single(message) = message.into_frame() else { continue; }; - let RawJsonRpcMessage::Response(v1::Response::Result { result, .. }) = message else { + let RawJsonRpcMessage::Response(agent_client_protocol::RawJsonRpcResponse::Result { + result, + .. + }) = message + else { continue; }; let initialize = v2::InitializeResponse::from_value("initialize", result)?; @@ -504,13 +508,13 @@ async fn assert_malformed_initialize_rejected(params: Map) -> Res let RawJsonRpcMessage::Response(response) = message else { continue; }; - let v1::Response::Error { error, .. } = response else { + let agent_client_protocol::RawJsonRpcResponse::Error { error, .. } = response else { panic!("malformed initialize should fail"); }; - assert_eq!(error.code, agent_client_protocol::ErrorCode::InvalidParams); + assert_eq!(error.code, -32602); let data = error .data - .as_ref() + .value() .and_then(|data| data.as_str()) .unwrap_or_default(); assert!(data.contains("protocolVersion"), "{error:?}"); @@ -1968,12 +1972,16 @@ async fn protocol_router_v2_only_rejects_v1_client() -> Result<(), Error> { let TransportFrame::Single(message) = message.into_frame() else { continue; }; - let RawJsonRpcMessage::Response(v1::Response::Error { error, .. }) = message else { + let RawJsonRpcMessage::Response(agent_client_protocol::RawJsonRpcResponse::Error { + error, + .. + }) = message + else { continue; }; let data = error .data - .as_ref() + .value() .and_then(|data| data.as_str()) .unwrap_or_default(); assert!( @@ -2715,7 +2723,11 @@ async fn protocol_router_routes_future_protocol_version_to_v2() -> Result<(), Er let TransportFrame::Single(message) = message.into_frame() else { continue; }; - let RawJsonRpcMessage::Response(v1::Response::Result { result, .. }) = message else { + let RawJsonRpcMessage::Response(agent_client_protocol::RawJsonRpcResponse::Result { + result, + .. + }) = message + else { continue; }; let initialize = v2::InitializeResponse::from_value("initialize", result)?; diff --git a/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs b/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs index 5df47ea8..e5feea27 100644 --- a/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs +++ b/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs @@ -95,8 +95,14 @@ async fn request( continue; }; match response { - v1::Response::Result { id, result } if id == request_id => break Ok(result), - v1::Response::Error { id, error } if id == request_id => break Err(error), + agent_client_protocol::RawJsonRpcResponse::Result { id, result } + if id == request_id => + { + break Ok(result); + } + agent_client_protocol::RawJsonRpcResponse::Error { id, error } if id == request_id => { + break Err(error.into_acp_error()); + } _ => {} } }; diff --git a/src/agent-client-protocol/tests/session_ordering.rs b/src/agent-client-protocol/tests/session_ordering.rs index a1850b8c..4a88b9ec 100644 --- a/src/agent-client-protocol/tests/session_ordering.rs +++ b/src/agent-client-protocol/tests/session_ordering.rs @@ -43,7 +43,7 @@ async fn initialize_raw_v2_proxy( .expect("proxy should accept initialization"); let Some(TransportFrame::Single(RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { id, result }, + agent_client_protocol::RawJsonRpcResponse::Result { id, result }, ))) = peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected the proxy initialize response"); @@ -465,7 +465,7 @@ async fn v2_proxy_session_start_installs_routing_before_later_batch_entry() { }; match message { RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { id, result }, + agent_client_protocol::RawJsonRpcResponse::Result { id, result }, ) => { assert_eq!(id, upstream_id); let response = v2::NewSessionResponse::from_value("session/new", result)?; @@ -624,7 +624,7 @@ async fn v2_proxy_fork_installs_response_id_routing_before_later_batch_entry() { }; match message { RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { id, result }, + agent_client_protocol::RawJsonRpcResponse::Result { id, result }, ) => { assert_eq!(id, upstream_id); let response = v2::ForkSessionResponse::from_value("session/fork", result)?; @@ -786,7 +786,7 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { )); let Some(TransportFrame::Single(RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { id, result }, + agent_client_protocol::RawJsonRpcResponse::Result { id, result }, ))) = peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected the resume response after replay"); From ae1cfac3377d38bd1f53829ff3a3d8aa827e9a64 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Mon, 28 Sep 2026 14:09:04 +0200 Subject: [PATCH 8/8] fix(acp): retain cancellation context and own connector operations Keep bounded HTTP response context until a terminal reply or teardown, and consistently bypass ordered POSTs for wrapped cancellation. Admit each connector operation as one protected task before construction, retaining ownership through backend shutdown and scoped cleanup. Add deterministic late-success, low-capacity, terminal-drain, and escaped-sender regressions. --- md/http-transport.md | 8 +- md/mcp-over-acp.md | 7 + src/agent-client-protocol-http/src/client.rs | 228 ++++++++++-- .../src/client_admission_tests.rs | 14 +- .../src/connection.rs | 9 - .../src/http_server.rs | 4 - .../src/protocol.rs | 17 +- src/agent-client-protocol/src/jsonrpc.rs | 38 ++ .../src/mcp_server/active_session.rs | 89 +++-- .../tests/mcp_connector_admission.rs | 336 ++++++++++++++++++ 10 files changed, 632 insertions(+), 118 deletions(-) create mode 100644 src/agent-client-protocol/tests/mcp_connector_admission.rs diff --git a/md/http-transport.md b/md/http-transport.md index b613875d..05316355 100644 --- a/md/http-transport.md +++ b/md/http-transport.md @@ -114,7 +114,13 @@ agent-client-protocol-http = { version = "...", features = ["client", "server"] `$/cancel_request` is connection-scoped. The HTTP transport does not apply `Acp-Session-Id` to cancellation notifications, and routes outgoing cancellation notifications over the connection stream rather than a session -stream. +stream. Cancellation is advisory: a `202 Accepted` for the cancellation POST +does not complete the original request. Pending response routing and its +bounded metadata charge stay in place until a terminal response, POST failure, +or connection teardown. The original request can still succeed after +cancellation, including a `session/new` or `session/fork` that opens a new +session stream. If the peer never sends a response, that request continues to +occupy pending-request capacity until the connection closes. WebSocket connections can carry cancellation at any point after the socket is open. With HTTP + SSE, cancellation can be sent after `initialize` completes and diff --git a/md/mcp-over-acp.md b/md/mcp-over-acp.md index 0dd66971..7e9c29ec 100644 --- a/md/mcp-over-acp.md +++ b/md/mcp-over-acp.md @@ -27,6 +27,13 @@ Use `McpServer::new_service` for a native service, or transport. The connector-based factory remains an explicit adapter for backends that require per-operation construction; stateless MCP does not require it. +A connector operation is admitted as one protected task before its factory is +called. That owner drives both the backend and response forwarding, then drops +the backend and joins scoped cleanup before releasing the logical ID or replying. +If the backend exits, already-accepted output is drained without waiting for +escaped sender handles; a valid queued terminal outcome is preserved, and later +notifications are not forwarded. + The rmcp integration's builder and `from_rmcp` use the reusable service path for ACP attachments. Each operation uses rmcp's direct, one-request transport without `initialize`. Its wrapper supervises rmcp handler futures through diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index f43ccc3c..3d0730f4 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -20,7 +20,7 @@ use thiserror::Error; use tracing::{debug, error, trace, warn}; use crate::protocol::{ - HEADER_CONNECTION_ID, HEADER_SESSION_ID, cancelled_request_id, is_initialize_request, + HEADER_CONNECTION_ID, HEADER_SESSION_ID, is_cancel_request_message, is_initialize_request, is_response_only_shape, method_for_message, method_requires_session_header, session_id_from_message, }; @@ -450,7 +450,6 @@ fn handle_completed_post( ) -> Result<(), AcpError> { let CompletedPost { pending_requests, - cancelled_requests, result, } = completed; if let Err(error) = result { @@ -458,9 +457,6 @@ fn handle_completed_post( error!("POST failed: {error}"); Err(AcpError::internal_error().data(format!("POST: {error}"))) } else { - for id in cancelled_requests { - state.cancel_pending_request(&id); - } Ok(()) } } @@ -514,11 +510,22 @@ fn is_response_only_frame(frame: &TransportFrame) -> bool { } fn is_cancellation_frame(frame: &TransportFrame) -> bool { - matches!( - frame, - TransportFrame::Single(RawJsonRpcMessage::Notification(message)) - if message.method.as_ref() == "$/cancel_request" - ) + match frame { + TransportFrame::Single(message) => is_cancel_request_message(message), + TransportFrame::Batch(batch) => { + let mut has_cancellation = false; + let only_control = batch.entries().all(|entry| match entry { + TransportBatchEntry::Message(RawJsonRpcMessage::Response(_)) => true, + TransportBatchEntry::Message(message) if is_cancel_request_message(message) => { + has_cancellation = true; + true + } + TransportBatchEntry::Malformed { .. } | TransportBatchEntry::Message(_) => false, + }); + only_control && has_cancellation + } + TransportFrame::Malformed { .. } => false, + } } enum HttpLoopEvent { @@ -864,7 +871,6 @@ struct ClientState { struct PendingPost { pending_requests: Vec<(RequestId, String)>, - cancelled_requests: Vec, response: BoxFuture<'static, Result<(), String>>, } @@ -872,13 +878,11 @@ impl PendingPost { fn into_completion(self, permit: Option) -> BoxFuture<'static, CompletedPost> { let Self { pending_requests, - cancelled_requests, response, } = self; async move { let completed = CompletedPost { pending_requests, - cancelled_requests, result: response.await, }; drop(permit); @@ -891,7 +895,6 @@ impl PendingPost { #[derive(Debug)] struct CompletedPost { pending_requests: Vec<(RequestId, String)>, - cancelled_requests: Vec, result: Result<(), String>, } @@ -1036,7 +1039,6 @@ impl ClientState { let pending_requests = pending_request_for_message(&msg) .into_iter() .collect::>(); - let cancelled_requests = cancelled_request_id(&msg).into_iter().collect(); self.check_pending_request_capacity(pending_requests.len())?; self.track_pending_requests(&pending_requests); @@ -1051,7 +1053,6 @@ impl ClientState { }; Ok(PendingPost { pending_requests, - cancelled_requests, response: response.boxed(), }) } @@ -1088,7 +1089,6 @@ impl ClientState { Ok(( PendingPost { pending_requests: bookkeeping.pending_requests, - cancelled_requests: bookkeeping.cancelled_requests, response: response.boxed(), }, session_ids, @@ -1160,20 +1160,6 @@ impl ClientState { method } - fn cancel_pending_request(&mut self, id: &RequestId) { - let Some(methods) = self.pending_requests.get_mut(id) else { - return; - }; - methods.pop_front(); - if let Some(leases) = self.pending_request_leases.get_mut(id) { - leases.pop_front(); - } - if methods.is_empty() { - self.pending_requests.remove(id); - self.pending_request_leases.remove(id); - } - } - fn register_session_streams( &mut self, session_ids: impl IntoIterator, @@ -1249,7 +1235,6 @@ impl ClientState { struct FrameBookkeeping { session_ids: Vec, pending_requests: Vec<(RequestId, String)>, - cancelled_requests: Vec, } impl FrameBookkeeping { @@ -1278,8 +1263,6 @@ impl FrameBookkeeping { if let Some(pending_request) = pending_request_for_message(message) { self.pending_requests.push(pending_request); } - self.cancelled_requests - .extend(cancelled_request_id(message)); Ok(()) } } @@ -1717,6 +1700,57 @@ mod tests { ); } + #[test] + fn cancel_ack_keeps_session_opening_context_until_terminal_response() { + for method in ["session/new", "session/fork"] { + let mut state = initialized_client_state(); + let id = RequestId::Number(7); + state.track_pending_requests(&[(id.clone(), method.into())]); + handle_completed_post( + &mut state, + CompletedPost { + pending_requests: Vec::new(), + result: Ok(()), + }, + ) + .unwrap(); + assert_eq!(state.pending_requests.get(&id).unwrap().len(), 1); + let response = single_frame(RawJsonRpcMessage::response( + id.clone(), + Ok(json!({"sessionId": "new-session"})), + )); + assert_eq!( + state.sessions_to_open_for_responses(&response), + ["new-session"] + ); + assert!(state.pending_requests.is_empty()); + assert!(state.sessions_to_open_for_responses(&response).is_empty()); + } + } + + #[test] + fn only_pure_cancellation_batches_bypass_ordered_posts() { + let cancel = || { + RawJsonRpcMessage::notification( + "_proxy/successor".into(), + json!({"method": "$/cancel_request", "params": {"requestId": 7}}), + ) + .unwrap() + }; + assert!(is_cancellation_frame(&single_frame(cancel()))); + let batch = |other| { + TransportFrame::Batch(TransportBatch::from_messages([cancel(), other]).unwrap()) + }; + assert!(is_cancellation_frame(&batch(cancel()))); + assert!(is_cancellation_frame(&batch(RawJsonRpcMessage::response( + RequestId::Number(7), + Ok(json!({})) + )))); + assert!(!is_cancellation_frame(&batch( + RawJsonRpcMessage::notification("custom/data".into(), json!({})).unwrap() + ))); + } + impl WsSink for RecordingWsSink { fn send( &mut self, @@ -2027,6 +2061,99 @@ mod tests { server.abort(); } + #[tokio::test] + async fn wrapped_cancellation_post_bypasses_blocked_ordered_post_without_session_header() { + let slow_started = Arc::new(Notify::new()); + let release_slow = Arc::new(Notify::new()); + let (cancel_tx, mut cancel_rx) = tokio::sync::mpsc::unbounded_channel(); + let app = Router::new().route( + "/acp", + post({ + let slow_started = slow_started.clone(); + let release_slow = release_slow.clone(); + move |headers: HeaderMap, body: String| { + let slow_started = slow_started.clone(); + let release_slow = release_slow.clone(); + let cancel_tx = cancel_tx.clone(); + async move { + let value: serde_json::Value = serde_json::from_str(&body).unwrap(); + match value["method"].as_str() { + Some("initialize") => initialize_response().await.into_response(), + Some("session/load") => { + slow_started.notify_one(); + release_slow.notified().await; + StatusCode::ACCEPTED.into_response() + } + Some("_proxy/successor") => { + cancel_tx.send((headers, value)).unwrap(); + StatusCode::ACCEPTED.into_response() + } + other => panic!("unexpected POST: {other:?}"), + } + } + } + }) + .get(pending_sse) + .delete(|| async { StatusCode::ACCEPTED }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let (mut caller, transport) = Channel::duplex(); + let transport = tokio::spawn(run( + HttpClient::new(format!("http://{addr}")).unwrap(), + transport, + )); + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::request("initialize".into(), json!({}), RequestId::Number(1)) + .unwrap(), + )) + .unwrap(); + timeout(Duration::from_secs(1), caller.rx.next()) + .await + .unwrap() + .unwrap(); + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::request( + "session/load".into(), + json!({"sessionId": "source"}), + RequestId::Number(2), + ) + .unwrap(), + )) + .unwrap(); + timeout(Duration::from_secs(1), slow_started.notified()) + .await + .unwrap(); + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::notification( + "_proxy/successor".into(), + json!({ + "method": "$/cancel_request", + "params": {"requestId": 2, "sessionId": "source"} + }), + ) + .unwrap(), + )) + .unwrap(); + let (headers, cancellation) = timeout(Duration::from_secs(1), cancel_rx.recv()) + .await + .expect("cancellation must bypass an in-flight session POST") + .unwrap(); + assert!(headers.get(HEADER_SESSION_ID).is_none()); + assert_eq!(cancellation["params"]["method"], "$/cancel_request"); + release_slow.notify_one(); + transport.abort(); + drop(caller); + server.abort(); + } + #[tokio::test] async fn http_preserves_batch_frames_across_post_and_sse() { let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel(); @@ -2300,6 +2427,37 @@ mod tests { .unwrap(); assert!(posted.is_array(), "outgoing batch must remain an array"); + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::notification("$/cancel_request".into(), json!({"requestId": 2})) + .unwrap(), + )) + .unwrap(); + let cancellation = timeout(Duration::from_secs(1), post_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(cancellation["method"], "$/cancel_request"); + assert_eq!(cancellation["params"]["requestId"], 2); + // Control POSTs are serialized with each other. Observing a second + // (unknown-ID, harmless) cancellation proves the client processed the + // first POST's 202 before the original successful response is emitted. + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::notification( + "$/cancel_request".into(), + json!({"requestId": 999}), + ) + .unwrap(), + )) + .unwrap(); + let barrier = timeout(Duration::from_secs(1), post_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(barrier["params"]["requestId"], 999); emit_response.notify_one(); let response = timeout(Duration::from_secs(1), caller.rx.next()) .await @@ -2897,7 +3055,6 @@ mod tests { let mut posts = PostQueues::default(); posts.ordered.push(PendingPost { pending_requests: vec![pending_request], - cancelled_requests: Vec::new(), response: async { Err("earlier post failed".to_string()) }.boxed(), }); @@ -2991,7 +3148,6 @@ mod tests { let mut posts = PostQueues::default(); posts.ordered.push(PendingPost { pending_requests: Vec::new(), - cancelled_requests: Vec::new(), response: async move { complete_earlier_post.notified().await; Ok(()) diff --git a/src/agent-client-protocol-http/src/client_admission_tests.rs b/src/agent-client-protocol-http/src/client_admission_tests.rs index 9e526fcd..58dfa374 100644 --- a/src/agent-client-protocol-http/src/client_admission_tests.rs +++ b/src/agent-client-protocol-http/src/client_admission_tests.rs @@ -134,7 +134,6 @@ async fn post_and_stream_counts_are_bounded_independently_of_frame_bytes() { check_post_capacity(&posts, 3, false).unwrap(); posts.ordered.push(PendingPost { pending_requests: Vec::new(), - cancelled_requests: Vec::new(), response: Box::pin(futures::future::pending()), }); } @@ -142,7 +141,6 @@ async fn post_and_stream_counts_are_bounded_independently_of_frame_bytes() { check_post_capacity(&posts, 3, true).unwrap(); posts.responses.push(PendingPost { pending_requests: Vec::new(), - cancelled_requests: Vec::new(), response: Box::pin(futures::future::pending()), }); assert_eq!(posts.len(), 3); @@ -168,7 +166,6 @@ async fn cancelled_post_releases_its_body_budget() { posts.ordered.push_budgeted( PendingPost { pending_requests: Vec::new(), - cancelled_requests: Vec::new(), response: Box::pin(futures::future::pending()), }, permit, @@ -212,7 +209,7 @@ async fn delivering_sse_preserves_admission_through_output_channel() { } #[tokio::test] -async fn pending_requests_hold_their_source_charge_until_response_or_cancel() { +async fn pending_requests_hold_their_source_charge_until_terminal_response() { let request = RawJsonRpcMessage::request( "test/request".to_string(), serde_json::json!({}), @@ -267,11 +264,18 @@ async fn pending_requests_hold_their_source_charge_until_response_or_cancel() { &mut state, CompletedPost { pending_requests: post.pending_requests, - cancelled_requests: post.cancelled_requests, result: Ok(()), }, ) .unwrap(); + assert!(state.check_pending_request_capacity(1).is_err()); + assert!(admission.try_admit(frame.clone()).is_err()); + assert_eq!( + state + .take_pending_request_method(&RequestId::Number(1)) + .as_deref(), + Some("test/request") + ); assert!(state.check_pending_request_capacity(1).is_ok()); assert!(admission.try_admit(frame).is_ok()); } diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index 07f0c0ba..cc1c7d68 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -189,15 +189,6 @@ impl Connection { } } - pub(crate) async fn cancel_pending_routes(&self, ids: &[RequestId]) { - if let OutboundTransport::Http(http) = &self.outbound_transport { - let mut pending = http.pending_routes.lock().await; - for id in ids { - take_pending_route(&mut pending, id); - } - } - } - #[cfg(test)] pub(crate) async fn ensure_session(&self, session_id: &str) { self.outbound_transport.ensure_session(session_id).await; diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index bda0274a..016a2937 100644 --- a/src/agent-client-protocol-http/src/http_server.rs +++ b/src/agent-client-protocol-http/src/http_server.rs @@ -145,7 +145,6 @@ pub(crate) async fn handle_post( let mut session_routes = Vec::new(); let mut pending_routes = Vec::new(); - let mut cancellations = Vec::new(); match &mut frame { TransportFrame::Single(message) => { let route = match prepare_message_route(message, session_id.as_deref()) { @@ -153,7 +152,6 @@ pub(crate) async fn handle_post( Err(error) => return (StatusCode::BAD_REQUEST, error).into_response(), }; collect_route(message, route, &mut session_routes, &mut pending_routes); - cancellations.extend(crate::protocol::cancelled_request_id(message)); trace!(connection_id = %connection_id, ?message, "POST → agent"); } TransportFrame::Batch(batch) => { @@ -166,7 +164,6 @@ pub(crate) async fn handle_post( Err(error) => return (StatusCode::BAD_REQUEST, error).into_response(), }; collect_route(message, route, &mut session_routes, &mut pending_routes); - cancellations.extend(crate::protocol::cancelled_request_id(message)); } trace!(connection_id = %connection_id, ?frame, "POST batch → agent"); } @@ -195,7 +192,6 @@ pub(crate) async fn handle_post( } drop(permit); inbound_slot.send(admitted); - connection.cancel_pending_routes(&cancellations).await; StatusCode::ACCEPTED.into_response() } diff --git a/src/agent-client-protocol-http/src/protocol.rs b/src/agent-client-protocol-http/src/protocol.rs index ff1cc495..5e2862f3 100644 --- a/src/agent-client-protocol-http/src/protocol.rs +++ b/src/agent-client-protocol-http/src/protocol.rs @@ -1,4 +1,4 @@ -use agent_client_protocol::{RawJsonRpcMessage, RawJsonRpcParams, schema::v1::RequestId}; +use agent_client_protocol::{RawJsonRpcMessage, RawJsonRpcParams}; pub(crate) const HEADER_CONNECTION_ID: &str = "acp-connection-id"; pub(crate) const HEADER_SESSION_ID: &str = "acp-session-id"; @@ -42,25 +42,12 @@ pub(crate) fn method_for_message(msg: &RawJsonRpcMessage) -> Option<&str> { } } -pub(crate) fn cancelled_request_id(msg: &RawJsonRpcMessage) -> Option { - let RawJsonRpcMessage::Notification(notification) = msg else { - return None; - }; - if notification.method.as_ref() != "$/cancel_request" { - return None; - } - let Some(RawJsonRpcParams::Object(params)) = notification.params.as_ref() else { - return None; - }; - serde_json::from_value(params.get("requestId")?.clone()).ok() -} - pub(crate) fn is_connection_scoped_protocol_message(msg: &RawJsonRpcMessage) -> bool { method_for_message(msg).is_some_and(|method| method.starts_with("$/")) || is_cancel_request_message(msg) } -fn is_cancel_request_message(msg: &RawJsonRpcMessage) -> bool { +pub(crate) fn is_cancel_request_message(msg: &RawJsonRpcMessage) -> bool { let RawJsonRpcMessage::Notification(notification) = msg else { return false; }; diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 280c4da1..76671ef8 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -7446,6 +7446,14 @@ pub struct FrameReceiver { _slots: async_channel::Sender<()>, } +#[cfg(feature = "unstable_mcp_over_acp")] +impl FrameReceiver { + /// Reject new output while retaining already-accepted frames for draining. + pub(crate) fn close(&self) { + self.rx.close(); + } +} + /// A slot is reserved before a frame enters the queue and returned at dequeue. /// Its drop also returns reservations abandoned by a cancelled send or sink. #[derive(Debug)] @@ -8037,6 +8045,36 @@ mod tests { .expect("dropping the last permit releases capacity"); } + #[cfg(feature = "unstable_mcp_over_acp")] + #[test] + fn receiver_close_drains_accepted_output_and_wakes_blocked_senders() { + let (source, mut destination) = Channel::duplex_with_limits(ConnectionLimits { + max_queued_frames: 1, + ..ConnectionLimits::default() + }); + let frame = capacity_frame(); + let escaped = source.tx.clone(); + source.tx.try_send(frame.clone()).unwrap(); + let mut blocked = Box::pin(escaped.send_frame(frame.clone())); + assert!(blocked.as_mut().now_or_never().is_none()); + + destination.rx.close(); + assert!( + blocked.now_or_never().unwrap().is_err(), + "receiver closure must wake a blocked producer" + ); + assert!(escaped.try_send(frame.clone()).is_err()); + let accepted = destination.rx.next().now_or_never().unwrap().unwrap(); + assert_eq!( + accepted.frame().to_json().unwrap(), + frame.to_json().unwrap() + ); + assert!( + destination.rx.next().now_or_never().unwrap().is_none(), + "draining must not wait for the escaped sender to be dropped" + ); + } + fn capacity_frame() -> TransportFrame { TransportFrame::Single( RawJsonRpcMessage::notification("capacity".into(), serde_json::json!({})).unwrap(), diff --git a/src/agent-client-protocol/src/mcp_server/active_session.rs b/src/agent-client-protocol/src/mcp_server/active_session.rs index e910434a..28d90f3e 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -425,30 +425,17 @@ where connection: connection.clone(), cleanup: Some(Arc::default()), }; - let backend = self.mcp_connect.connect(cleanup_connection.clone()); + let connector = self.mcp_connect.clone(); let connection_for_task = connection.clone(); let cancellation = responder.cancellation(); - let (mut client, server) = Channel::duplex(); - // Keep the operation admitted until its backend has actually stopped. - let (backend_stop_tx, backend_stop_rx) = oneshot::channel::<()>(); - let (backend_done_tx, mut backend_done_rx) = oneshot::channel(); - let spawn_result = connection.spawn_protected(async move { - // Own (not merely borrow) the future so cancellation drops its - // backend before the completion acknowledgement is published. - let run = Box::pin(backend.connect_to(server)); - let outcome = match future::select(run, backend_stop_rx).await { - Either::Left((result, _)) => result, - Either::Right((_, _)) => Ok(()), - }; - drop(backend_done_tx.send(outcome)); - Ok(()) - }); - if let Err(error) = spawn_result { - drop(guard); - responder.respond_with_error(error)?; - return Ok(Handled::Yes); - } - let spawn_result = connection.spawn_protected(async move { + // Admission is atomic: this one protected task owns construction, + // execution, forwarding, and cleanup. A rejected spawn cannot start + // a backend or leave half of the operation running. + connection.spawn_protected(async move { + let backend = connector.connect(cleanup_connection.clone()); + let (mut client, server) = Channel::duplex(); + let mut backend = Some(Box::pin(backend.connect_to(server))); + let mut backend_error = None; let inner_id = RequestId::Str(request_id.0.to_string()); let is_discovery = method == "server/discover"; let process = async { @@ -462,7 +449,31 @@ where .send_frame(TransportFrame::Single(raw)) .await .map_err(crate::Error::into_internal_error)?; - while let Some(budgeted) = client.rx.next().await { + loop { + let message = match backend.as_mut() { + Some(run) => match future::select(client.rx.next(), run).await { + Either::Left((message, _)) => message, + Either::Right((result, receive)) => { + drop(receive); + backend.take(); + backend_error = result.err(); + // Drain accepted output after backend exit, + // without waiting for an escaped sender to + // close or accepting any later output. + client.rx.close(); + continue; + } + }, + None => client.rx.next().await, + }; + let Some(budgeted) = message else { + return Err(backend_error.take().unwrap_or_else(|| { + crate::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend closed without a response", + ) + })); + }; let (frame, _permit) = budgeted.into_parts(); let TransportFrame::Single(message) = frame else { return Err(crate::Error::new( @@ -524,10 +535,6 @@ where } } } - Err(crate::Error::new( - MCP_BACKEND_FAILURE, - "MCP backend closed without a response", - )) }; let result = cancellation .run_until_cancelled(async { @@ -541,28 +548,16 @@ where .await; }; futures::pin_mut!(stop); - let work = async { - match future::select(process, stop).await { - Either::Left((result, _)) => result, - Either::Right(((), _)) => Err(crate::Error::request_cancelled()), - } - }; - futures::pin_mut!(work); - match future::select(work, &mut backend_done_rx).await { + match future::select(process, stop).await { Either::Left((result, _)) => result, - Either::Right((Ok(Err(error)), _)) => Err(error), - // The backend can finish immediately after queueing its - // reply. Drain the channel before calling that an EOF. - Either::Right((Ok(Ok(())) | Err(_), work)) => work.await, + Either::Right(((), _)) => Err(crate::Error::request_cancelled()), } }) .await; - // Revoking the channel stops any late output. A cancellation is only - // caller-visible now; cleanup and ID release happen after backend exit. - drop(backend_stop_tx); - // The receiver can have already completed in the race above. Polling - // it again then returns immediately; otherwise this joins cleanup. - drop(backend_done_rx.await); + // Drop the owned driver and revoke its channels before joining + // scoped tool cleanup, releasing the logical ID, or replying. + drop(backend); + drop(client); cleanup_connection.wait_cleanup().await; drop(guard); let response = send_outcome::(responder, result, is_discovery); @@ -570,9 +565,7 @@ where tracing::debug!(?error, "cannot send request-scoped MCP response"); } Ok(()) - }); - // A failed spawn drops its responder and backend stop sender with the task. - spawn_result?; + })?; Ok(Handled::Yes) } } diff --git a/src/agent-client-protocol/tests/mcp_connector_admission.rs b/src/agent-client-protocol/tests/mcp_connector_admission.rs new file mode 100644 index 00000000..70155ff6 --- /dev/null +++ b/src/agent-client-protocol/tests/mcp_connector_admission.rs @@ -0,0 +1,336 @@ +#![cfg(feature = "unstable_mcp_over_acp")] + +use std::{ + future::pending, + sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, ByteStreams, Channel, Client, ConnectTo, ConnectionLimits, ConnectionTo, DynConnectTo, + Error, FrameSender, RawJsonRpcMessage, Responder, RunWithConnectionTo, TransportFrame, + mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, + role, + schema::v1, +}; +use futures::StreamExt as _; +use serde_json::{Map, Value, json}; +use tokio::{ + io::{AsyncBufReadExt, AsyncWriteExt, BufReader, duplex}, + sync::oneshot, +}; +use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; + +const TIMEOUT: Duration = Duration::from_secs(10); + +#[derive(Default)] +struct Probes { + factory: AtomicUsize, + backend_started: AtomicUsize, + backend_dropped: AtomicUsize, + escaped_senders: Mutex>, + notifications: Mutex>, +} + +#[derive(Clone, Copy)] +enum BackendBehavior { + WireReply, + ReplyThenError, + ExitWithoutReply, +} + +struct Connector(Arc, BackendBehavior); + +impl McpServerConnect for Connector { + fn name(&self) -> String { + "capacity-probe".into() + } + + fn connect(&self, context: McpConnectionTo) -> DynConnectTo { + assert!(context.request_id().is_some()); + self.0.factory.fetch_add(1, Ordering::SeqCst); + DynConnectTo::new(Backend(self.0.clone(), self.1)) + } +} + +struct Backend(Arc, BackendBehavior); + +struct BackendDrop(Arc); + +impl Drop for BackendDrop { + fn drop(&mut self) { + self.0.backend_dropped.fetch_add(1, Ordering::SeqCst); + } +} + +impl ConnectTo for Backend { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + self.0.backend_started.fetch_add(1, Ordering::SeqCst); + let _drop = BackendDrop(self.0.clone()); + if !matches!(self.1, BackendBehavior::WireReply) { + let (mut channel, driver) = client.into_channel_and_future(); + let work = async { + let frame = channel.rx.next().await.expect("MCP request"); + let TransportFrame::Single(RawJsonRpcMessage::Request(request)) = + frame.into_frame() + else { + panic!("expected one request"); + }; + self.0 + .escaped_senders + .lock() + .unwrap() + .push(channel.tx.clone()); + if matches!(self.1, BackendBehavior::ExitWithoutReply) { + return Ok(()); + } + for message in [ + RawJsonRpcMessage::notification( + "notifications/progress".into(), + json!({"progressToken":1, "progress":1, "marker":"before"}), + )?, + RawJsonRpcMessage::response(request.id, Ok(json!({"admitted":true}))), + RawJsonRpcMessage::notification( + "notifications/progress".into(), + json!({"progressToken":1, "progress":2, "marker":"after"}), + )?, + ] { + channel + .tx + .try_send(TransportFrame::Single(message)) + .map_err(Error::into_internal_error)?; + } + Err(Error::internal_error().data("driver failed after accepted output")) + }; + let (driver, result) = tokio::join!(driver, work); + driver?; + return result; + } + let (sdk_output, peer_input) = duplex(4096); + let (peer_output, sdk_input) = duplex(4096); + let transport = ByteStreams::new(sdk_output.compat_write(), sdk_input.compat()); + let peer = async move { + let mut reader = BufReader::new(peer_input); + let mut line = String::new(); + if reader.read_line(&mut line).await.expect("read MCP request") == 0 { + // Rejected admissions close the backend before sending anything. + return; + } + let request: Value = serde_json::from_str(&line).expect("valid MCP request"); + assert_eq!( + request["params"]["_meta"]["io.modelcontextprotocol/protocolVersion"], + "2026-07-28" + ); + let response = json!({ + "jsonrpc": "2.0", + "id": request["id"], + "result": {"admitted": true} + }); + let mut output = peer_output; + output + .write_all(format!("{response}\n").as_bytes()) + .await + .expect("write MCP response"); + output.shutdown().await.expect("close MCP backend output"); + }; + let (result, ()) = tokio::join!(client.connect_to(transport), peer); + result + } +} + +struct NullRun; + +impl RunWithConnectionTo for NullRun { + async fn run_with_connection_to(self, _connection: ConnectionTo) -> Result<(), Error> { + pending().await + } +} + +// Unlike a raw Channel, this uses ConnectTo's default adapter. Only the +// provider endpoint is limited; the agent must have its own default pool. +struct DefaultCapacityAgent(Channel); + +impl ConnectTo for DefaultCapacityAgent { + async fn connect_to(self, agent: impl ConnectTo) -> Result<(), Error> { + ConnectTo::::connect_to(self.0, agent).await + } +} + +fn params() -> Map { + serde_json::from_value(json!({ + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + "progressToken": 1 + } + })) + .unwrap() +} + +// Occupy slots only after newSession has completed. The agent waits on `start` +// so no MCP request can race the provider's admission setup. +async fn scenario( + fillers: usize, + requests: usize, + behavior: BackendBehavior, +) -> ( + Vec<(Result, usize)>, + Arc, +) { + tokio::time::timeout(TIMEOUT, async move { + let probes = Arc::new(Probes::default()); + let (provider_channel, agent_channel) = Channel::duplex_with_limits(ConnectionLimits { + max_queued_frames: 8, + ..ConnectionLimits::default() + }); + let (start_tx, start_rx) = oneshot::channel::<()>(); + let start_rx = Mutex::new(Some(start_rx)); + let (done_tx, done_rx) = oneshot::channel(); + let done_tx = Mutex::new(Some(done_tx)); + let agent_probes = probes.clone(); + let notification_probes = probes.clone(); + let agent = Agent.builder().on_receive_request( + async move |request: v1::NewSessionRequest, + responder: Responder, + connection: ConnectionTo| { + let [v1::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected one native MCP server") + }; + let server_id = server.server_id.clone(); + responder.respond(v1::NewSessionResponse::new(v1::SessionId::new( + "capacity-session", + )))?; + let start = start_rx.lock().unwrap().take().expect("one session"); + let done = done_tx.lock().unwrap().take().expect("one session"); + let sender = connection.clone(); + let observed = agent_probes.clone(); + connection.spawn(async move { + start.await.map_err(Error::into_internal_error)?; + let mut responses = Vec::new(); + for _ in 0..requests { + // Reuse the logical ID after the first request completes. + let response = sender + .send_request( + v1::MessageMcpRequest::new( + server_id.clone(), + "reused-id", + "admission/probe", + ) + .params(params()), + ) + .block_task() + .await; + responses.push((response, observed.backend_dropped.load(Ordering::SeqCst))); + } + drop(done.send(responses)); + Ok(()) + }) + }, + agent_client_protocol::on_receive_request!(), + ); + let agent = agent.on_receive_notification( + async move |notification: v1::MessageMcpNotification, _cx| { + assert_eq!(notification.request_id, v1::McpRequestId::new("reused-id")); + notification_probes.notifications.lock().unwrap().push( + notification.params.as_ref().unwrap()["marker"] + .as_str() + .unwrap() + .to_owned(), + ); + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ); + let connector = Connector(probes.clone(), behavior); + let test = Client + .builder() + .connect_with(provider_channel, async move |connection| { + let filler_connection = connection.clone(); + connection + .build_session_cwd()? + .with_mcp_server(McpServer::::new(connector, NullRun))? + .block_task() + .run_until(async |_session| { + for _ in 0..fillers { + filler_connection + .spawn(async { pending::>().await })?; + } + start_tx.send(()).expect("agent still waiting"); + let responses = done_rx.await.map_err(Error::into_internal_error)?; + Ok(responses) + }) + .await + }); + let (result, agent_result) = + tokio::join!(test, DefaultCapacityAgent(agent_channel).connect_to(agent)); + agent_result.expect("agent connection"); + (result.expect("provider connection"), probes) + }) + .await + .expect("MCP capacity scenario timed out") +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn one_free_slot_completes_connector_and_recovers_for_reused_id() { + let (responses, probes) = scenario(7, 2, BackendBehavior::WireReply).await; + for (index, (response, dropped_at_response)) in responses.into_iter().enumerate() { + match response.expect("one slot must admit the MCP request") { + v1::MessageMcpResponse::Result { result, .. } => { + assert_eq!(result, json!({"admitted": true})); + } + other => panic!("expected MCP result, got {other:?}"), + } + assert_eq!( + dropped_at_response, + index + 1, + "backend must stop before the logical response is observed" + ); + } + assert_eq!(probes.factory.load(Ordering::SeqCst), 2); + assert_eq!(probes.backend_started.load(Ordering::SeqCst), 2); + assert_eq!(probes.backend_dropped.load(Ordering::SeqCst), 2); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn zero_free_slots_rejects_before_factory_or_backend_setup() { + let (responses, probes) = scenario(8, 1, BackendBehavior::WireReply).await; + let [(response, _)] = <[_; 1]>::try_from(responses).expect("one request"); + let error = response.expect_err("all live slots are occupied"); + assert!(error.to_string().contains("live task capacity"), "{error}"); + assert_eq!(probes.factory.load(Ordering::SeqCst), 0); + assert_eq!(probes.backend_started.load(Ordering::SeqCst), 0); + assert_eq!(probes.backend_dropped.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn connector_drains_terminal_output_before_driver_failure_and_rejects_late_output() { + let (responses, probes) = scenario(7, 1, BackendBehavior::ReplyThenError).await; + let [(response, dropped)] = <[_; 1]>::try_from(responses).unwrap(); + let v1::MessageMcpResponse::Result { result, .. } = response.unwrap() else { + panic!("accepted terminal result must survive backend exit"); + }; + assert_eq!(result, json!({"admitted":true})); + assert_eq!(dropped, 1); + assert_eq!(*probes.notifications.lock().unwrap(), ["before"]); + let escaped = probes.escaped_senders.lock().unwrap(); + assert_eq!(escaped.len(), 1); + assert!(escaped[0].is_closed()); +} + +#[tokio::test] +async fn connector_exit_without_response_does_not_wait_for_escaped_sender() { + let (responses, probes) = scenario(7, 1, BackendBehavior::ExitWithoutReply).await; + let [(response, dropped)] = <[_; 1]>::try_from(responses).unwrap(); + let error = response.expect_err("completed backend did not produce a terminal outcome"); + assert_eq!( + i32::from(error.code), + agent_client_protocol::mcp_server::MCP_BACKEND_FAILURE + ); + assert_eq!(dropped, 1); + let escaped = probes.escaped_senders.lock().unwrap(); + assert_eq!(escaped.len(), 1); + assert!(escaped[0].is_closed()); +}