From ee9bf7ed5207e4f247931101ef7ae0b216241183 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Thu, 24 Sep 2026 13:51:33 +0200 Subject: [PATCH 1/5] 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/5] 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/5] 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/5] 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/5] 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()); +}