From 1f3e69df03013b81df0caf5c8298f01fc666d2b8 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Thu, 24 Sep 2026 12:27:31 +0200 Subject: [PATCH] fix(mcp): checkpoint transport lifecycle and isolation work --- Cargo.lock | 1 + md/SUMMARY.md | 1 + md/mcp-bridge.md | 51 +- md/mcp-over-acp.md | 105 ++++ md/protocol.md | 42 +- .../Cargo.toml | 1 + .../tests/mcp_over_acp_polyfill.rs | 199 ++++++- .../tests/mcp_over_acp_polyfill_v2.rs | 214 ++++++- .../CHANGELOG.md | 13 + .../src/mcp_over_acp/actor.rs | 35 +- .../src/mcp_over_acp/http.rs | 262 +++++++-- .../src/mcp_over_acp/mod.rs | 46 +- src/agent-client-protocol-rmcp/Cargo.toml | 5 + .../examples/native_mcp_over_acp.rs | 197 +++++++ src/agent-client-protocol/CHANGELOG.md | 9 + src/agent-client-protocol/src/lib.rs | 3 + src/agent-client-protocol/src/mcp_client.rs | 511 ++++++++++++++++ .../src/mcp_server/active_session.rs | 151 ++++- .../tests/mcp_connection_lifecycle.rs | 545 ++++++++++++++++++ .../tests/native_mcp_consumer.rs | 319 ++++++++++ 20 files changed, 2567 insertions(+), 143 deletions(-) create mode 100644 md/mcp-over-acp.md create mode 100644 src/agent-client-protocol-rmcp/examples/native_mcp_over_acp.rs create mode 100644 src/agent-client-protocol/src/mcp_client.rs create mode 100644 src/agent-client-protocol/tests/mcp_connection_lifecycle.rs create mode 100644 src/agent-client-protocol/tests/native_mcp_consumer.rs diff --git a/Cargo.lock b/Cargo.lock index 9e2bdbf0..cc0e9373 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -46,6 +46,7 @@ dependencies = [ "futures", "futures-concurrency", "regex", + "reqwest", "rmcp", "rustc-hash", "schemars 1.2.2", diff --git a/md/SUMMARY.md b/md/SUMMARY.md index e9fd0e84..716cfd28 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..4fd33e6f 100644 --- a/md/mcp-bridge.md +++ b/md/mcp-bridge.md @@ -1,5 +1,12 @@ # MCP-over-ACP Compatibility Bridge +**Draft checkpoint:** The stateful adapter described here targets older MCP +semantics. It is not an implementation of MCP 2026-07-28, which removes +initialization, protocol sessions, GET, and DELETE. The intended MCP-over-ACP +transport will target that stateless revision only; retaining this session mode +for backwards compatibility is not a goal. See the +[modernization audit](https://agentclientprotocol.com/rfds/mcp-over-acp#modernization-audit-mcp-2026-07-28). + `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. @@ -93,12 +100,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. +3. Opens a native connection when an HTTP MCP client initializes a logical + session, sending `mcp/connect` with that server ID toward the provider. + Independent HTTP sessions receive independent native connections. 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. +5. Sends an `mcp/disconnect` request when the logical HTTP MCP session closes + and removes that connection from the bridge. The listening endpoint remains + available for other sessions. Enable the polyfill crate's `unstable_session_fork` feature when adapting fork requests. Stable v1 setup includes `session/new`, `session/load`, and @@ -120,8 +129,15 @@ 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 MCP POST requests, SSE GET streams, and session DELETE +requests at `/`, retaining JSON-RPC batch frames and correlating each POST with +its response. + +The adapter uses stateful Streamable HTTP. A successful MCP initialization +returns an `MCP-Session-Id` header. Clients must send that header on subsequent +POST, GET, and DELETE requests; unknown or closed sessions return HTTP 404. +Clients that previously ignored session headers must retain the returned ID. +An individual POST response or GET stream closing does not end the session. ```rust,ignore let bridge = McpOverAcpPolyfill::http(); @@ -132,11 +148,18 @@ implement resumable SSE event IDs. ## 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. +The listener and the logical MCP connections have different lifetimes. Endpoint +creation alone does not open an MCP connection. Each HTTP MCP session receives +its own native `connectionId` from `mcp/connect`, so initialization, request IDs, +and server-originated messages cannot cross between clients. Disconnecting one +session leaves its siblings and the cached endpoint usable. + +Deleting an HTTP MCP session stops its local transport and sends +`mcp/disconnect` for that session's native connection. Request failures use the +corresponding request's error path; notifications are never answered with +synthetic errors. Closing the parent ACP connection drops its listeners and +session tasks; a disconnect exchange is not possible after that transport is +gone. A reverse `mcp/message` request for an unknown `connectionId` receives `Invalid params`. A reverse notification for an unknown connection is ignored, @@ -144,3 +167,9 @@ as required for JSON-RPC notifications. The polyfill does not infer or store ACP session IDs. Association is carried by the declared `serverId` and the resulting active `connectionId`. + +Known checkpoint gaps: aborting HTTP initialization before receiving its +response can leave an unadvertised session until the listener stops. There is +no idle timeout, and DELETE during pending forward/reverse requests still +needs regression coverage. These are reasons to keep the checkpoint in draft, +not features to preserve in the stateless replacement. diff --git a/md/mcp-over-acp.md b/md/mcp-over-acp.md new file mode 100644 index 00000000..40f7299a --- /dev/null +++ b/md/mcp-over-acp.md @@ -0,0 +1,105 @@ +# Native MCP-over-ACP + +**Draft checkpoint:** This chapter describes the in-progress connection-oriented +implementation, not conformance with MCP 2026-07-28. The stabilization target is +that stateless MCP revision only, with no legacy compatibility requirement. +Its initialization and connect/disconnect API below will need replacement; see +the [modernization audit](https://agentclientprotocol.com/rfds/mcp-over-acp#modernization-audit-mcp-2026-07-28). + +MCP-over-ACP lets an ACP client or proxy provide an MCP server through its +existing ACP connection. A native agent can use that server without starting +another process, opening an HTTP endpoint, or introducing a conductor. + +Enable `unstable_mcp_over_acp` on `agent-client-protocol`. Add +`unstable_protocol_v2` for draft-v2 connections. The wire types come from the +shared ACP schema; the transport remains unstable. + +## Providing a server + +The `mcp_server::McpServer` APIs attach servers to session setup requests using +`McpServer::Acp` declarations. The separate `agent-client-protocol-rmcp` crate +can build a server from tools or an `rmcp` service without making `rmcp` a +dependency of the core SDK. + +There are three distinct identifiers: + +| Identifier | Meaning | +| --- | --- | +| `serverId` | Provider-generated identity for the declared MCP server | +| `connectionId` | Provider-generated identity for one active connection to that server | +| Outer JSON-RPC `id` | Identity of one ACP request, including an `mcp/message` request | + +A server can accept multiple connections. Each has independent MCP +initialization, pending requests, and shutdown. Providers must not reuse a +server ID for different servers visible on the same ACP connection. + +Servers are ready as soon as their declarations are published. An agent can +connect and run MCP initialization before returning the ACP session ID. +Providers must not wait for the session setup response before serving MCP. + +## Consuming a server + +The core `mcp_client` module supplies `McpOverAcp`, a server transport for an +ordinary MCP client. Its version-specific connection helpers open a native +connection to the declared server and route bidirectional MCP traffic over +the ACP channel. No HTTP adapter is involved. + +The consuming ACP agent holds a `ConnectionTo` (or its v2 counterpart): +the ACP client is providing the MCP server. The returned transport implements +`ConnectTo`, so it can be used by the SDK's MCP client role +or connected to an external MCP implementation through a byte-stream adapter. + +The helper opens the transport, not the MCP protocol session. The MCP client +still performs its normal `initialize` / `notifications/initialized` handshake. +It can then list and call tools, while also handling server-originated requests +and notifications. + +Do connection setup and MCP work outside the ACP dispatch loop, for example +in a connection-spawned task. Awaiting a peer response inside an ACP message +handler can block the very messages needed to complete that operation. See +[Ordered Application Dispatch](./ordered-application-dispatch.md). + +## Closing connections + +Use the helper's awaited close operation when shutdown completion matters. +Dropping a native consumer schedules best-effort disconnect; it cannot report +whether the provider acknowledged cleanup. + +The provider acknowledges `mcp/disconnect` after stopping that connection's +relay and server work. Requests already dispatched to the child are completed +or failed; this checkpoint still needs explicit failure of requests queued in +the relay at shutdown and direct handler-drop coverage. Other MCP connections +to the same server, and the containing ACP connection, stay usable. + +ACP transport closure drops its MCP connection-scoped work. No disconnect +exchange is possible once the ACP transport is gone. + +## Runnable direct example + +From the repository root: + +```sh +cargo run -p agent-client-protocol-rmcp \ + --example native_mcp_over_acp \ + --features native_mcp_example +``` + +The example connects an ACP client directly to an ACP agent, attaches an +`rmcp` server on the client side, and uses a real MCP client on the agent side. +It exercises the normal MCP handshake and tool traffic without a conductor, +an HTTP listener, or a subprocess. + +## Compatibility + +An agent advertises native support with +`agentCapabilities.mcpCapabilities.acp: true` in v1, or +`capabilities.session.mcp.acp: {}` in draft v2. Do not advertise support unless +the agent can consume the transport. + +For an HTTP-capable agent without native support, use the explicit +[MCP-over-ACP compatibility bridge](./mcp-bridge.md). Its listening endpoint +can be shared, but each logical HTTP MCP session has its own native connection. + +See the [protocol reference](./protocol.md#native-mcp-over-acp) for exact +wire envelopes and the [RFD](https://agentclientprotocol.com/rfds/mcp-over-acp) +for the protocol design. diff --git a/md/protocol.md b/md/protocol.md index d30d6747..30fe3a02 100644 --- a/md/protocol.md +++ b/md/protocol.md @@ -5,6 +5,12 @@ conductor and the opt-in native MCP-over-ACP transport exposed by the shared ACP schema. The proxy methods are provisional SDK extensions. MCP-over-ACP is also unstable and is available only with the `unstable_mcp_over_acp` feature. +The MCP wire methods below describe the current **connection-oriented +checkpoint**, not the intended MCP 2026-07-28-only design. Modernization will +replace the old lifecycle with stateless requests and request-scoped +notifications/cancellation; see the +[RFD audit](https://agentclientprotocol.com/rfds/mcp-over-acp#modernization-audit-mcp-2026-07-28). + ## Method Summary | Method | JSON-RPC shape | Purpose | @@ -63,8 +69,9 @@ inner message. 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: +(`session/new` and `session/resume`, plus `session/load` in v1 and the opt-in +`session/fork` in either version). Its wire shape contains a human-readable +name and an opaque server identifier: ```json { @@ -81,9 +88,15 @@ for multiple visible servers on the same ACP connection. The high-level 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. +`agentCapabilities.mcpCapabilities.acp: true` in v1 or +`capabilities.session.mcp.acp: {}` in draft v2. The v2 capability is an optional +object: omission or `null` means support is not advertised. The `mcp/*` wire +envelopes below are the same in both versions. + +See [Native MCP-over-ACP](./mcp-over-acp.md) for the consumer helper and a direct +client/agent example. 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` @@ -111,6 +124,13 @@ connection ID: The server ID selects what to connect to; the connection ID selects that particular running connection. All subsequent messages use the connection ID. +Each connect request creates a separate MCP connection, even when the server +ID is reused. MCP initialization still happens inside that connection through +`mcp/message`. + +The provider must be ready when it publishes the server declaration. An agent +may connect and initialize its MCP servers before returning the ACP session +ID; routing cannot depend on waiting for the session setup response. ### `mcp/message` @@ -135,7 +155,10 @@ is bidirectional because MCP clients and servers can both issue requests: 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. +error directly, without another JSON-RPC envelope. The inner `params` field +accepts an object or `null`; omission and `null` both mean no parameters. +Positional arrays are not supported. ACP `_meta` alongside `connectionId` is +separate from MCP metadata inside the inner parameters or result. ### `mcp/disconnect` @@ -161,8 +184,15 @@ A successful disconnect returns an empty result: } ``` +The provider stops the connection's relay and server work before acknowledging +disconnect. Outstanding requests complete or fail, and further messages +cannot use the closed connection. Sibling MCP connections and the parent ACP +connection remain usable. An MCP server failure is contained to that connection; +closing the ACP connection releases all of its MCP connections. + ## Related Documentation +- [MCP-over-ACP RFD](https://agentclientprotocol.com/rfds/mcp-over-acp) - [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/Cargo.toml b/src/agent-client-protocol-conductor/Cargo.toml index 7b49d9a4..f23145a5 100644 --- a/src/agent-client-protocol-conductor/Cargo.toml +++ b/src/agent-client-protocol-conductor/Cargo.toml @@ -42,6 +42,7 @@ agent-client-protocol-test.workspace = true yopo.workspace = true expect-test.workspace = true regex.workspace = true +reqwest.workspace = true rmcp = { workspace = true, features = [ "client", "server", 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..120f8c09 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,10 +6,11 @@ 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, - ResumeSessionResponse, SessionCapabilities, SessionResumeCapabilities, + AgentCapabilities, ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, + DisconnectMcpResponse, InitializeRequest, InitializeResponse, LoadSessionRequest, + LoadSessionResponse, McpCapabilities, McpServer, McpServerAcp, MessageMcpRequest, + NewSessionRequest, NewSessionResponse, ResumeSessionRequest, ResumeSessionResponse, + SessionCapabilities, SessionResumeCapabilities, }; use agent_client_protocol::{Agent, Client, Conductor, ConnectTo, Proxy}; use agent_client_protocol_conductor::{ConductorImpl, ProxiesAndAgent}; @@ -36,6 +37,8 @@ struct SetupRequest { #[derive(Default)] struct ObservedRequests { setup: Mutex>, + native_messages: Mutex>, + disconnect_count: AtomicUsize, } impl ObservedRequests { @@ -57,6 +60,7 @@ struct RecordingAgent { struct NativeMcpProvider { connect_count: Arc, + observed: Arc, } impl ConnectTo for NativeMcpProvider { @@ -64,6 +68,8 @@ impl ConnectTo for NativeMcpProvider { self, client: impl ConnectTo, ) -> Result<(), agent_client_protocol::Error> { + let messages = Arc::clone(&self.observed); + let disconnected = Arc::clone(&self.observed); Proxy .builder() .name("native-mcp-provider") @@ -71,8 +77,41 @@ impl ConnectTo for NativeMcpProvider { Agent, async move |request: ConnectMcpRequest, 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")) + let index = self.connect_count.fetch_add(1, Ordering::SeqCst); + responder.respond(ConnectMcpResponse::new(format!("test-connection-{index}"))) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request_from( + Agent, + async move |request: MessageMcpRequest, responder, _cx| { + messages + .native_messages + .lock() + .unwrap() + .push((request.connection_id.to_string(), request.method.clone())); + let result = match request.method.as_str() { + "initialize" => serde_json::json!({ + "protocolVersion": "2025-06-18", + "capabilities": { "tools": {} }, + "serverInfo": { "name": "v1-test", "version": "1" } + }), + "tools/list" => serde_json::json!({ "tools": [] }), + _ => { + 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!(), + ) + .on_receive_request_from( + Agent, + async move |_request: DisconnectMcpRequest, responder, _cx| { + disconnected.disconnect_count.fetch_add(1, Ordering::SeqCst); + responder.respond(DisconnectMcpResponse::new()) }, agent_client_protocol::on_receive_request!(), ) @@ -163,6 +202,7 @@ async fn run_with_polyfill( agent_client_protocol::ConnectionTo, ) -> Result<(), agent_client_protocol::Error>, ) -> Result<(), agent_client_protocol::Error> { + let observed = Arc::clone(&agent.observed); drop( tracing_subscriber::fmt() .with_env_filter(tracing_subscriber::EnvFilter::from_default_env()) @@ -185,6 +225,7 @@ async fn run_with_polyfill( ProxiesAndAgent::new(agent) .proxy(NativeMcpProvider { connect_count: provider_connect_count, + observed, }) .proxy(McpOverAcpPolyfill::http()), ) @@ -245,8 +286,8 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ .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" + 0, + "creating and reusing the listener must not open a logical MCP session" ); assert_eq!(setup.len(), 3); assert_eq!(setup[0].method, SetupMethod::New); @@ -282,6 +323,148 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ Ok(()) } +#[tokio::test] +async fn v1_http_sessions_open_lazily_and_disconnect_independently() +-> Result<(), agent_client_protocol::Error> { + let observed = Arc::new(ObservedRequests::default()); + let connects = Arc::new(AtomicUsize::new(0)); + run_with_polyfill( + RecordingAgent { + capabilities: agent_capabilities(McpCapabilities::new().http(true)), + observed: Arc::clone(&observed), + }, + Arc::clone(&connects), + async |connection| { + recv(connection.send_request(InitializeRequest::new(ProtocolVersion::V1))).await?; + recv(connection.send_request( + NewSessionRequest::new(PathBuf::from("/tmp")).mcp_servers(vec![native_server()]), + )) + .await?; + let endpoint = { + let setup = observed.setup.lock().unwrap(); + let McpServer::Http(server) = &setup[0].mcp_servers[0] else { + panic!("expected HTTP adaptation"); + }; + server.url.clone() + }; + assert_eq!(connects.load(Ordering::SeqCst), 0); + let http = reqwest::Client::new(); + let init = serde_json::json!({ + "jsonrpc": "2.0", "id": 1, "method": "initialize", + "params": { + "protocolVersion": "2025-06-18", "capabilities": {}, + "clientInfo": { "name": "v1-client", "version": "1" } + } + }); + let (a, b) = tokio::join!( + http.post(&endpoint).json(&init).send(), + http.post(&endpoint).json(&init).send(), + ); + let a = a.unwrap(); + let b = b.unwrap(); + let a_id = a + .headers() + .get("mcp-session-id") + .unwrap() + .to_str() + .unwrap() + .to_owned(); + let b_id = b + .headers() + .get("mcp-session-id") + .unwrap() + .to_str() + .unwrap() + .to_owned(); + assert_ne!(a_id, b_id); + assert_eq!(connects.load(Ordering::SeqCst), 2); + assert!(a.text().await.unwrap().contains("\"result\"")); + assert!(b.text().await.unwrap().contains("\"result\"")); + let tool = serde_json::json!({ + "jsonrpc": "2.0", "id": 1, "method": "tools/list", "params": {} + }); + let (a, b) = tokio::join!( + http.post(&endpoint) + .header("mcp-session-id", &a_id) + .json(&tool) + .send(), + http.post(&endpoint) + .header("mcp-session-id", &b_id) + .json(&tool) + .send(), + ); + assert!(a.unwrap().text().await.unwrap().contains("\"tools\":[]")); + assert!(b.unwrap().text().await.unwrap().contains("\"tools\":[]")); + assert_eq!( + http.delete(&endpoint) + .header("mcp-session-id", &a_id) + .send() + .await + .unwrap() + .status(), + reqwest::StatusCode::ACCEPTED + ); + assert_eq!(observed.disconnect_count.load(Ordering::SeqCst), 1); + assert_eq!( + http.post(&endpoint) + .header("mcp-session-id", &a_id) + .json(&tool) + .send() + .await + .unwrap() + .status(), + reqwest::StatusCode::NOT_FOUND + ); + assert!( + http.post(&endpoint) + .header("mcp-session-id", &b_id) + .json(&tool) + .send() + .await + .unwrap() + .text() + .await + .unwrap() + .contains("\"tools\":[]") + ); + let c = http.post(&endpoint).json(&init).send().await.unwrap(); + let c_id = c + .headers() + .get("mcp-session-id") + .unwrap() + .to_str() + .unwrap() + .to_owned(); + assert_ne!(c_id, a_id); + assert_eq!(connects.load(Ordering::SeqCst), 3); + for id in [&b_id, &c_id] { + assert_eq!( + http.delete(&endpoint) + .header("mcp-session-id", id) + .send() + .await + .unwrap() + .status(), + reqwest::StatusCode::ACCEPTED + ); + } + assert_eq!(observed.disconnect_count.load(Ordering::SeqCst), 3); + assert_eq!( + observed + .native_messages + .lock() + .unwrap() + .iter() + .filter(|(_, method)| method == "initialize") + .count(), + 3 + ); + Ok(()) + }, + ) + .await +} + #[tokio::test] async fn native_downstream_keeps_capability_and_declaration_unchanged() -> Result<(), agent_client_protocol::Error> { 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..34cd7fef 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 @@ -44,6 +44,7 @@ struct SetupRequest { #[derive(Default)] struct ObservedRequests { setup: Mutex>, + native_messages: Mutex>, } impl ObservedRequests { @@ -112,6 +113,7 @@ struct NativeMcpProvider { request_methods: Arc>>, notification_methods: Arc>>, disconnect_count: Arc, + observed: Arc, } impl ConnectTo for NativeMcpProvider { @@ -122,6 +124,7 @@ impl ConnectTo for NativeMcpProvider { let request_methods = Arc::clone(&self.request_methods); let notification_methods = Arc::clone(&self.notification_methods); let disconnect_count = Arc::clone(&self.disconnect_count); + let observed = Arc::clone(&self.observed); Proxy .v2() @@ -130,14 +133,21 @@ impl ConnectTo for NativeMcpProvider { Agent, async move |request: v2::ConnectMcpRequest, 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")) + let index = self.connect_count.fetch_add(1, Ordering::SeqCst); + responder.respond(v2::ConnectMcpResponse::new(format!( + "v2-test-connection-{index}" + ))) }, agent_client_protocol::on_receive_request!(), ) .on_receive_request_from( Agent, async move |request: v2::MessageMcpRequest, responder, _cx| { + observed + .native_messages + .lock() + .unwrap() + .push((request.connection_id.to_string(), request.method.clone())); request_methods .lock() .expect("request method mutex should not be poisoned") @@ -247,6 +257,7 @@ async fn run_with_polyfill( provider_disconnect_count: Arc, editor_task: impl AsyncFnOnce(V2ConnectionTo) -> Result<(), agent_client_protocol::Error>, ) -> Result<(), agent_client_protocol::Error> { + let observed = Arc::clone(&agent.observed); let (editor_out, conductor_in) = duplex(4096); let (conductor_out, editor_in) = duplex(4096); let transport = @@ -264,6 +275,7 @@ async fn run_with_polyfill( request_methods: provider_request_methods, notification_methods: provider_notification_methods, disconnect_count: provider_disconnect_count, + observed, }) .proxy(McpOverAcpPolyfill::http()), ) @@ -421,6 +433,204 @@ async fn http_downstream_adapts_v2_capabilities_and_only_transforms_native_serve Ok(()) } +async fn http_call( + client: &reqwest::Client, + endpoint: &str, + session: Option<&str>, + method: &str, +) -> (reqwest::StatusCode, Option, serde_json::Value) { + let params = if method == "initialize" { + serde_json::json!({ + "protocolVersion": "2025-06-18", + "capabilities": {}, + "clientInfo": { "name": "session-isolation-test", "version": "1" } + }) + } else { + serde_json::json!({}) + }; + let mut request = client.post(endpoint).json(&serde_json::json!({ + "jsonrpc": "2.0", "id": 1, "method": method, "params": params + })); + if let Some(session) = session { + request = request.header("mcp-session-id", session); + } + let response = request.send().await.expect("HTTP MCP POST"); + let status = response.status(); + let id = response + .headers() + .get("mcp-session-id") + .map(|value| value.to_str().expect("ASCII session ID").to_owned()); + let body = response.text().await.expect("HTTP MCP response body"); + let payload = body + .lines() + .find_map(|line| line.strip_prefix("data: ")) + .map(|line| serde_json::from_str(line).expect("JSON-RPC SSE data")) + .unwrap_or(serde_json::Value::Null); + (status, id, payload) +} + +#[tokio::test] +async fn concurrent_http_sessions_isolate_ids_and_survive_individual_delete() +-> 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), + }; + let connects = Arc::new(AtomicUsize::new(0)); + let disconnects = Arc::new(AtomicUsize::new(0)); + run_with_polyfill( + agent, + Arc::clone(&connects), + Arc::new(Mutex::new(Vec::new())), + Arc::new(Mutex::new(Vec::new())), + Arc::clone(&disconnects), + async |connection| { + connection + .send_request(initialize_request()) + .block_task() + .await?; + connection + .send_request( + v2::NewSessionRequest::new(PathBuf::from("/tmp")) + .mcp_servers(vec![native_server()]), + ) + .block_task() + .await?; + let endpoint = { + let setup = observed.setup.lock().unwrap(); + let v2::McpServer::Http(server) = &setup[0].mcp_servers[0] else { + panic!("expected adapted HTTP server"); + }; + server.url.clone() + }; + let http = reqwest::Client::new(); + let (a, b) = tokio::join!( + http_call(&http, &endpoint, None, "initialize"), + http_call(&http, &endpoint, None, "initialize"), + ); + assert_eq!(a.0, reqwest::StatusCode::OK); + assert_eq!(b.0, reqwest::StatusCode::OK); + assert!( + a.2.get("result").is_some(), + "first initialization: {:?}", + a.2 + ); + assert!( + b.2.get("result").is_some(), + "second initialization: {:?}", + b.2 + ); + let a_id = a.1.expect("first MCP-Session-Id"); + let b_id = b.1.expect("second MCP-Session-Id"); + assert_ne!(a_id, b_id); + assert_eq!(connects.load(Ordering::SeqCst), 2); + + // The same JSON-RPC ID on two sessions must not serialize or + // cross-deliver replies. Both session headers remain stable. + let (a, b) = tokio::join!( + http_call(&http, &endpoint, Some(&a_id), "tools/list"), + http_call(&http, &endpoint, Some(&b_id), "tools/list"), + ); + for (status, header, response) in [a, b] { + assert_eq!(status, reqwest::StatusCode::OK); + assert!(header.is_none(), "session ID belongs only on initialize"); + assert_eq!(response["id"], 1); + assert_eq!(response["result"]["tools"], serde_json::json!([])); + } + let unknown = http + .post(&endpoint) + .header("mcp-session-id", "not-a-session") + .json(&serde_json::json!({"jsonrpc":"2.0","id":1,"method":"tools/list"})) + .send() + .await + .unwrap(); + assert_eq!(unknown.status(), reqwest::StatusCode::NOT_FOUND); + let absent = http.get(&endpoint).send().await.unwrap(); + assert_eq!(absent.status(), reqwest::StatusCode::BAD_REQUEST); + + let deleted = http + .delete(&endpoint) + .header("mcp-session-id", &a_id) + .send() + .await + .unwrap(); + assert_eq!(deleted.status(), reqwest::StatusCode::ACCEPTED); + assert_eq!( + disconnects.load(Ordering::SeqCst), + 1, + "DELETE must wait for native disconnect" + ); + assert_eq!( + http.delete(&endpoint) + .header("mcp-session-id", &a_id) + .send() + .await + .unwrap() + .status(), + reqwest::StatusCode::NOT_FOUND + ); + assert_eq!( + http_call(&http, &endpoint, Some(&a_id), "tools/list") + .await + .0, + reqwest::StatusCode::NOT_FOUND + ); + assert_eq!( + http_call(&http, &endpoint, Some(&b_id), "tools/list") + .await + .2["result"]["tools"], + serde_json::json!([]) + ); + + let c = http_call(&http, &endpoint, None, "initialize").await; + assert!(c.2.get("result").is_some()); + let c_id = c.1.expect("new session after DELETE"); + assert_ne!(c_id, a_id); + assert_eq!(connects.load(Ordering::SeqCst), 3); + assert_eq!( + http_call(&http, &endpoint, Some(&c_id), "tools/list") + .await + .2["result"]["tools"], + serde_json::json!([]) + ); + for id in [&b_id, &c_id] { + assert_eq!( + http.delete(&endpoint) + .header("mcp-session-id", id) + .send() + .await + .unwrap() + .status(), + reqwest::StatusCode::ACCEPTED + ); + } + assert_eq!(disconnects.load(Ordering::SeqCst), 3); + let routes = observed.native_messages.lock().unwrap(); + let mut per_connection = std::collections::BTreeMap::new(); + for (id, method) in routes.iter() { + per_connection + .entry(id.clone()) + .or_insert_with(Vec::new) + .push(method.as_str()); + } + assert_eq!(per_connection.len(), 3); + assert!(per_connection.values().all(|methods| { + methods.first() == Some(&"initialize") + && methods + .iter() + .filter(|method| **method == "tools/list") + .count() + >= 1 + })); + Ok(()) + }, + ) + .await +} + #[tokio::test] async fn native_v2_downstream_keeps_capability_and_declarations_unchanged() -> Result<(), agent_client_protocol::Error> { diff --git a/src/agent-client-protocol-polyfill/CHANGELOG.md b/src/agent-client-protocol-polyfill/CHANGELOG.md index 77cb9487..5ecd0f9f 100644 --- a/src/agent-client-protocol-polyfill/CHANGELOG.md +++ b/src/agent-client-protocol-polyfill/CHANGELOG.md @@ -7,6 +7,19 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Fixed + +- Isolate logical HTTP MCP sessions behind a shared bridge endpoint. Each + session opens its own native MCP-over-ACP connection on initialization and + disconnects independently, leaving sibling sessions and the endpoint usable. + +### Changed + +- The HTTP bridge uses stateful Streamable HTTP: clients must retain the + `MCP-Session-Id` returned by initialization and include it in subsequent + POST, GET, and DELETE requests. Endpoint creation no longer eagerly opens an + MCP connection. + ## [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/src/mcp_over_acp/actor.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs index 908b6b87..f9670434 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs @@ -1,5 +1,9 @@ use agent_client_protocol::{ConnectTo, Dispatch, DynConnectTo, role::mcp}; -use futures::{SinkExt as _, StreamExt as _, channel::mpsc}; +use futures::{ + SinkExt as _, StreamExt as _, + channel::{mpsc, oneshot}, + future::Either, +}; use tracing::info; use super::BridgeMessage; @@ -31,7 +35,11 @@ impl BridgeConnectionActor { } } - pub async fn run(self, connection_id: String) -> Result<(), agent_client_protocol::Error> { + pub async fn run( + self, + connection_id: String, + disconnected_tx: oneshot::Sender>, + ) -> Result<(), agent_client_protocol::Error> { info!(connection_id, "MCP bridge connected"); let Self { @@ -61,15 +69,32 @@ impl BridgeConnectionActor { ) .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)?; + loop { + let next = to_mcp_client_rx.next(); + let closed = mcp_connection_to_client.incoming_closed(); + match futures::future::select(Box::pin(next), Box::pin(closed)).await { + Either::Left((Some(message), _)) => { + mcp_connection_to_client.send_proxied_message(message)?; + } + Either::Left((None, _)) | Either::Right(((), _)) => break, + } + } + // The runner still holds the sender until it receives + // Disconnected. Reject already queued reverse calls explicitly. + while let Ok(message) = to_mcp_client_rx.try_recv() { + if let Dispatch::Request(_, responder) = message { + drop(responder.respond_with_internal_error("HTTP MCP session closed")); + } } Ok(()) }) .await; bridge_tx - .send(BridgeMessage::Disconnected { connection_id }) + .send(BridgeMessage::Disconnected { + connection_id, + disconnected_tx, + }) .await .map_err(|_| agent_client_protocol::Error::internal_error())?; 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..18065111 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 @@ -11,20 +11,24 @@ use agent_client_protocol::{ use axum::{ Router, extract::State, - http::StatusCode, + http::{HeaderMap, HeaderValue, StatusCode}, response::{IntoResponse, Response, Sse}, routing::post, }; -use futures::{SinkExt, StreamExt as _, channel::mpsc, future::Either, stream::Stream}; -use futures_concurrency::future::FutureExt as _; +use futures::{ + SinkExt, StreamExt as _, + channel::{mpsc, oneshot}, + future::Either, + stream::Stream, +}; use futures_concurrency::stream::StreamExt as _; use rustc_hash::FxHashMap; use std::{ collections::{HashMap, VecDeque}, - pin::pin, sync::Arc, }; use tokio::net::TcpListener; +use tokio::sync::Mutex; use super::{BridgeConnection, BridgeMessage, actor::BridgeConnectionActor}; @@ -32,37 +36,29 @@ use super::{BridgeConnection, BridgeMessage, actor::BridgeConnectionActor}; pub async fn run_http_listener( tcp_listener: TcpListener, server_id: String, - mut bridge_tx: mpsc::Sender, + bridge_tx: mpsc::Sender, ) -> Result<(), agent_client_protocol::Error> { - let (to_mcp_client_tx, to_mcp_client_rx) = mpsc::channel(128); - - 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), - }) + let state = Arc::new(ListenerState { + server_id, + bridge_tx, + sessions: Mutex::new(HashMap::new()), + }); + let app = Router::new() + .route( + "/", + post(session_post).get(session_get).delete(session_delete), + ) + .with_state(state); + axum::serve(tcp_listener, app) .await - .map_err(|_| agent_client_protocol::Error::internal_error())?; - - Ok(()) + .map_err(agent_client_protocol::util::internal_error) } -/// A component that receives HTTP requests/responses using the HTTP transport -/// defined by the MCP protocol. +/// Each logical HTTP session has its own raw-frame router and native connection. struct HttpMcpBridge { - listener: tokio::net::TcpListener, -} - -impl HttpMcpBridge { - /// Creates a new HTTP-MCP bridge from an existing TCP listener. - fn new(listener: tokio::net::TcpListener) -> Self { - Self { listener } - } + client_channel: Channel, + server_channel: Channel, + registration_rx: mpsc::UnboundedReceiver, } impl ConnectTo for HttpMcpBridge { @@ -71,7 +67,7 @@ impl ConnectTo for HttpMcpBridge { 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 { + match futures::future::select(Box::pin(client.connect_to(channel)), serve_self).await { Either::Left((result, _)) | Either::Right((result, _)) => result, } } @@ -85,8 +81,10 @@ impl ConnectTo for HttpMcpBridge { where Self: Sized, { - let (channel_a, channel_b) = Channel::duplex(); - (channel_a, Box::pin(run(self.listener, channel_b))) + ( + self.client_channel, + Box::pin(RunningServer::new().run(self.server_channel, self.registration_rx)), + ) } } @@ -108,36 +106,173 @@ impl IntoResponse for HttpError { } } -/// 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) +const SESSION_HEADER: &str = "mcp-session-id"; + +struct ListenerState { + server_id: String, + bridge_tx: mpsc::Sender, + sessions: Mutex>, +} + +struct Session { + state: Arc, + disconnected: oneshot::Receiver>, +} + +impl BridgeState { + fn close(&self) { + drop(self.registration_tx.unbounded_send(HttpMessage::Close)); + } +} + +fn session_id(headers: &HeaderMap) -> Option<&str> { + headers + .get(SESSION_HEADER) + .and_then(|value| value.to_str().ok()) +} + +async fn session_get(State(state): State>, headers: HeaderMap) -> Response { + let Some(id) = session_id(&headers) else { + return StatusCode::BAD_REQUEST.into_response(); + }; + let session = state + .sessions + .lock() + .await + .get(id) + .map(|session| session.state.clone()); + let Some(session) = session else { + return StatusCode::NOT_FOUND.into_response(); + }; + match handle_get(State(session)).await { + Ok(response) => response.into_response(), + Err(error) => error.into_response(), + } +} + +async fn session_delete(State(state): State>, headers: HeaderMap) -> Response { + let Some(id) = session_id(&headers) else { + return StatusCode::BAD_REQUEST.into_response(); + }; + if let Some(session) = state.sessions.lock().await.remove(id) { + session.state.close(); + match session.disconnected.await { + Ok(Ok(())) => StatusCode::ACCEPTED.into_response(), + Ok(Err(error)) => { + tracing::warn!(?error, "native MCP disconnect failed"); + StatusCode::BAD_GATEWAY.into_response() + } + Err(_) => StatusCode::BAD_GATEWAY.into_response(), + } + } else { + StatusCode::NOT_FOUND.into_response() + } +} + +async fn session_post( + State(state): State>, + headers: HeaderMap, + body: String, +) -> Response { + if let Some(id) = session_id(&headers) { + let session = state + .sessions + .lock() .await - .map_err(agent_client_protocol::util::internal_error) + .get(id) + .map(|session| session.state.clone()); + return match session { + Some(session) => handle_post(State(session), body) + .await + .unwrap_or_else(IntoResponse::into_response), + None => StatusCode::NOT_FOUND.into_response(), + }; + } + + // Only an initialize request without a session ID can establish a new + // logical session. In particular, malformed frames never open connections. + let TransportFrame::Single(RawJsonRpcMessage::Request(request)) = + TransportFrame::parse_json(&body) + else { + return StatusCode::BAD_REQUEST.into_response(); + }; + if request.method.as_ref() != "initialize" { + return StatusCode::BAD_REQUEST.into_response(); + } + + let id = uuid::Uuid::new_v4().to_string(); + let (registration_tx, registration_rx) = mpsc::unbounded(); + let (client_channel, server_channel) = Channel::duplex(); + let session = Arc::new(BridgeState { registration_tx }); + let (to_mcp_client_tx, to_mcp_client_rx) = mpsc::channel(128); + let (disconnected_tx, disconnected) = oneshot::channel(); + let actor = BridgeConnectionActor::new( + HttpMcpBridge { + client_channel, + server_channel, + registration_rx, + }, + state.bridge_tx.clone(), + to_mcp_client_rx, + ); + state.sessions.lock().await.insert( + id.clone(), + Session { + state: session.clone(), + disconnected, + }, + ); + let mut bridge_tx = state.bridge_tx.clone(); + if bridge_tx + .send(BridgeMessage::ConnectionReceived { + server_id: state.server_id.clone(), + actor, + connection: BridgeConnection::new(to_mcp_client_tx), + disconnected_tx, + }) + .await + .is_err() + { + state.sessions.lock().await.remove(&id); + session.close(); + return StatusCode::SERVICE_UNAVAILABLE.into_response(); } - .race(RunningServer::new().run(channel, registration_rx)) - .await + + let (tx, mut rx) = mpsc::unbounded(); + if session + .registration_tx + .unbounded_send(HttpMessage::Request { + http_request_id: uuid::Uuid::new_v4(), + request, + response_tx: tx, + }) + .is_err() + { + state.sessions.lock().await.remove(&id); + session.close(); + return StatusCode::SERVICE_UNAVAILABLE.into_response(); + } + // Don't hand out a session ID unless initialization actually succeeded. + let Some(response) = rx.next().await else { + state.sessions.lock().await.remove(&id); + session.close(); + return StatusCode::SERVICE_UNAVAILABLE.into_response(); + }; + let success = matches!( + &response, + TransportFrame::Single(RawJsonRpcMessage::Response(RpcResponse::Result { .. })) + ); + if !success { + state.sessions.lock().await.remove(&id); + session.close(); + return immediate_sse_response(response); + } + let mut http_response = immediate_sse_response(response); + http_response.headers_mut().insert( + SESSION_HEADER, + HeaderValue::from_str(&id).expect("UUID is a valid HTTP header value"), + ); + http_response } /// The state we pass to our POST/GET handlers. @@ -150,6 +285,7 @@ struct BridgeState { #[derive(Debug)] #[allow(dead_code)] enum HttpMessage { + Close, /// A JSON-RPC request (has an id, expects a response via the channel). Request { http_request_id: uuid::Uuid, @@ -220,6 +356,9 @@ impl RunningServer { match message { MultiplexMessage::FromHttpToChannel(http_message) => { + if matches!(http_message, HttpMessage::Close) { + return Ok(()); + } self.handle_http_message(http_message, &mut channel.tx)?; } MultiplexMessage::FromChannelToHttp(message) => { @@ -241,6 +380,7 @@ impl RunningServer { channel_tx: &mut mpsc::UnboundedSender, ) -> Result<(), agent_client_protocol::Error> { match message { + HttpMessage::Close => return Ok(()), HttpMessage::Request { http_request_id, request, 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..471dab63 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 @@ -59,6 +59,7 @@ pub(crate) enum BridgeMessage { server_id: String, actor: BridgeConnectionActor, connection: BridgeConnection, + disconnected_tx: oneshot::Sender>, }, /// A native MCP connection ID was received; spawn the actor and store its sender. @@ -67,6 +68,7 @@ pub(crate) enum BridgeMessage { connection_id: String, actor: BridgeConnectionActor, connection: BridgeConnection, + disconnected_tx: oneshot::Sender>, }, /// Opening a native MCP connection failed. @@ -88,7 +90,10 @@ pub(crate) enum BridgeMessage { ServerToClientNotification { notification: NativeMcpMessage }, /// The local MCP bridge disconnected. - Disconnected { connection_id: String }, + Disconnected { + connection_id: String, + disconnected_tx: oneshot::Sender>, + }, } /// Connection handle for sending messages to an MCP client via a bridge. @@ -452,10 +457,6 @@ impl BridgeListeners { self.listeners.insert(server_id, listener); Ok(declaration) } - - fn remove(&mut self, server_id: &str) { - self.listeners.remove(server_id); - } } #[derive(Debug)] @@ -532,13 +533,13 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { server_id, actor, connection: bridge, + disconnected_tx, } => { let Some(protocol) = self.protocol else { warn!( server_id, "cannot open MCP bridge before ACP initialization" ); - self.listeners.remove(&server_id); continue; }; let request = protocol.connect_request(server_id.clone())?; @@ -553,6 +554,7 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { connection_id, actor, connection: bridge, + disconnected_tx, }, Err(error) => { warn!(?error, "invalid response to mcp/connect"); @@ -577,16 +579,20 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { connection_id, actor, connection: bridge, + disconnected_tx, } => { self.bridge_connections.insert( connection_id.clone(), ActiveBridgeConnection { server_id, bridge }, ); - connection.spawn(actor.run(connection_id))?; + connection.spawn(actor.run(connection_id, disconnected_tx))?; } BridgeMessage::ConnectionFailed { server_id } => { - self.listeners.remove(&server_id); + warn!( + server_id, + "MCP session connection failed; listener remains available" + ); } BridgeMessage::ClientToServer { @@ -741,13 +747,15 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { } } - BridgeMessage::Disconnected { connection_id } => { + BridgeMessage::Disconnected { + connection_id, + disconnected_tx, + } => { 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); - + debug!(server_id = %active.server_id, connection_id, "closing MCP session"); let Some(protocol) = self.protocol else { debug!("could not disconnect MCP bridge before ACP initialization"); continue; @@ -756,18 +764,10 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { 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"); - } - } + let result = result.and_then(|response| { + protocol.validate_disconnect_response(response) + }); + drop(disconnected_tx.send(result)); Ok(()) }); if let Err(error) = scheduled { diff --git a/src/agent-client-protocol-rmcp/Cargo.toml b/src/agent-client-protocol-rmcp/Cargo.toml index 4952b98d..03c20725 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"] +native_mcp_example = ["unstable_mcp_over_acp", "rmcp/client"] [[example]] name = "with_mcp_server" required-features = ["unstable_mcp_over_acp"] +[[example]] +name = "native_mcp_over_acp" +required-features = ["native_mcp_example"] + [dependencies] agent-client-protocol = { workspace = true, features = ["schemars"] } futures.workspace = true diff --git a/src/agent-client-protocol-rmcp/examples/native_mcp_over_acp.rs b/src/agent-client-protocol-rmcp/examples/native_mcp_over_acp.rs new file mode 100644 index 00000000..f7cf2d9b --- /dev/null +++ b/src/agent-client-protocol-rmcp/examples/native_mcp_over_acp.rs @@ -0,0 +1,197 @@ +//! Direct in-memory ACP agent/client with a native MCP server and an rmcp client. +//! +//! Run: `cargo run -p agent-client-protocol-rmcp --example native_mcp_over_acp --features native_mcp_example` +//! The provider below explicitly forwards native ACP messages to its local +//! rmcp server to keep this example self-contained. The *agent* only needs +//! `McpOverAcp` to consume the declaration; no conductor or HTTP bridge runs. + +use std::sync::{Arc, Mutex}; + +use agent_client_protocol::{ + Agent, ByteStreams, Channel, Client, ConnectionTo, Dispatch, Error, JsonRpcResponse, Responder, + UntypedMessage, + mcp_client::McpOverAcp, + mcp_server::McpServer, + role, + schema::v1::{ + ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, DisconnectMcpResponse, + McpConnectionId, McpServerAcp, McpServerAcpId, MessageMcpNotification, MessageMcpRequest, + MessageMcpResponse, + }, +}; +use agent_client_protocol_rmcp::McpServerExt; +use futures::{StreamExt, channel::mpsc}; +use rmcp::{ + ErrorData as McpError, ServerHandler, + handler::server::{router::tool::ToolRouter, wrapper::Parameters}, + model::*, + tool, tool_handler, tool_router, +}; +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; + +#[derive(Debug, Serialize, Deserialize, schemars::JsonSchema)] +struct EchoParams { + message: String, +} + +#[derive(Clone, Debug)] +struct EchoServer { + #[allow(dead_code)] + tool_router: ToolRouter, +} + +#[tool_router] +impl EchoServer { + #[tool(description = "Return the supplied message")] + async fn echo( + &self, + Parameters(params): Parameters, + ) -> Result { + Ok(CallToolResult::success(vec![ContentBlock::text( + params.message, + )])) + } +} + +#[allow(unknown_lints, clippy::unused_async_trait_impl)] +#[tool_handler] +impl ServerHandler for EchoServer { + fn get_info(&self) -> ServerInfo { + ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) + .with_server_info(Implementation::new("native-example", "1")) + } +} + +fn provider( + active: Arc>>>, + mut incoming: mpsc::UnboundedReceiver, +) -> impl agent_client_protocol::ConnectTo { + let active_requests = active.clone(); + let active_notifications = active.clone(); + Client + .builder() + .on_receive_request( + async |request: ConnectMcpRequest, responder: Responder, _| { + if request.server_id.0.as_ref() != "echo" { + return responder.respond_with_error(Error::invalid_params()); + } + responder.respond(ConnectMcpResponse::new(McpConnectionId::new( + "echo-connection", + ))) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: MessageMcpRequest, responder: Responder, _| { + if request.connection_id.0.as_ref() != "echo-connection" { + return responder.respond_with_error(Error::invalid_params()); + } + let responder = responder.wrap_params(|method, result| { + result.and_then(|value: Value| MessageMcpResponse::from_value(method, value)) + }); + let active = active_requests.lock().unwrap(); + let Some(sender) = active.as_ref() else { + return responder.respond_with_error(Error::internal_error()); + }; + sender + .unbounded_send(Dispatch::Request( + UntypedMessage { + method: request.method, + params: request.params.map_or(Value::Null, Value::Object), + }, + responder, + )) + .map_err(Error::into_internal_error) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_notification( + async move |notification: MessageMcpNotification, _| { + let active = active_notifications.lock().unwrap(); + let Some(sender) = active.as_ref() else { + return Ok(()); + }; + sender + .unbounded_send(Dispatch::Notification(UntypedMessage { + method: notification.method, + params: notification.params.map_or(Value::Null, Value::Object), + })) + .map_err(Error::into_internal_error) + }, + agent_client_protocol::on_receive_notification!(), + ) + .on_receive_request( + async move |request: DisconnectMcpRequest, + responder: Responder, + _| { + assert_eq!(request.connection_id.0.as_ref(), "echo-connection"); + active.lock().unwrap().take(); + responder.respond(DisconnectMcpResponse::new()) + }, + agent_client_protocol::on_receive_request!(), + ) + .with_spawned(move |_acp: ConnectionTo| async move { + let server: McpServer = + McpServer::from_rmcp("echo", || EchoServer { + tool_router: EchoServer::tool_router(), + }); + role::mcp::Client + .builder() + .connect_with(server, async |mcp| { + while let Some(message) = incoming.next().await { + mcp.send_proxied_message_to(role::mcp::Server, message)?; + } + Ok(()) + }) + .await + }) +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + let declaration = McpServerAcp::new("echo", McpServerAcpId::new("echo")); + let (acp_agent, acp_client) = Channel::duplex(); + let (outgoing, incoming) = mpsc::unbounded(); + let active = Arc::new(Mutex::new(Some(outgoing))); + + let agent = Agent.builder().connect_with(acp_agent, async |cx| { + let (transport, close) = McpOverAcp::connect_v1(&cx, declaration.server_id.clone()).await?; + let (sdk_stream, rmcp_stream) = tokio::io::duplex(8192); + let (sdk_read, sdk_write) = tokio::io::split(sdk_stream); + let (rmcp_read, rmcp_write) = tokio::io::split(rmcp_stream); + let sdk_mcp = agent_client_protocol::ConnectTo::::connect_to( + transport, + ByteStreams::new(sdk_write.compat_write(), sdk_read.compat()), + ); + let rmcp_client = async move { + let service = rmcp::serve_client((), (rmcp_read, rmcp_write)) + .await + .map_err(Error::into_internal_error)?; + let tools = service + .peer() + .list_all_tools() + .await + .map_err(Error::into_internal_error)?; + assert!(tools.iter().any(|tool| tool.name == "echo")); + println!( + "MCP tools over direct ACP: {:?}", + tools.iter().map(|tool| &tool.name).collect::>() + ); + close.close().await?; + drop(service); + Ok::<_, Error>(()) + }; + futures::try_join!(sdk_mcp, rmcp_client)?; + Ok(()) + }); + futures::try_join!( + agent, + agent_client_protocol::ConnectTo::::connect_to( + provider(active, incoming), + acp_client + ) + )?; + Ok(()) +} diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index fa0324ee..14e45136 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -10,6 +10,15 @@ servers and independently enabled unstable protocol features remain available. Existing users of `default-features = false` who need the previous JSON Schema or typed MCP tool APIs should add `features = ["schemars"]`. +- Add a runtime-agnostic native MCP-over-ACP consumer transport for ACP agents, + supporting v1 and draft v2, bidirectional MCP traffic, and awaited close + without a conductor or HTTP bridge. + +### Fixed + +- Stop native MCP-over-ACP relay and server tasks before acknowledging + `mcp/disconnect`. Contain individual MCP connection failures and clean up + connection-scoped work without terminating sibling MCP connections or ACP. ## [2.2.0](https://github.com/agentclientprotocol/rust-sdk/compare/v2.1.0...v2.2.0) - 2026-09-18 diff --git a/src/agent-client-protocol/src/lib.rs b/src/agent-client-protocol/src/lib.rs index 0a543aed..71c71341 100644 --- a/src/agent-client-protocol/src/lib.rs +++ b/src/agent-client-protocol/src/lib.rs @@ -131,6 +131,9 @@ pub mod component; pub mod concepts; /// JSON-RPC connection and handler infrastructure mod jsonrpc; +/// Native MCP-over-ACP transport for agents consuming client-provided servers. +#[cfg(feature = "unstable_mcp_over_acp")] +pub mod mcp_client; /// Runtime-agnostic MCP server support, including optional attachment to ACP sessions. pub mod mcp_server; /// Role types for ACP connections diff --git a/src/agent-client-protocol/src/mcp_client.rs b/src/agent-client-protocol/src/mcp_client.rs new file mode 100644 index 00000000..0344532e --- /dev/null +++ b/src/agent-client-protocol/src/mcp_client.rs @@ -0,0 +1,511 @@ +//! Consume client-provided native MCP servers from an ACP agent connection. +//! +//! [`McpOverAcp::connect_v1`] opens the declared server before returning a +//! standard `ConnectTo` transport. Keep the transport alive +//! while using it; its [`McpOverAcpClose`] handle permits awaited shutdown from +//! inside the client's connection callback. Dropping the last close handle +//! unregisters routing and schedules a best-effort disconnect. + +use std::{ + collections::HashMap, + marker::PhantomData, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, AtomicU64, Ordering}, + }, +}; + +use futures::{ + StreamExt, + channel::{mpsc, oneshot}, +}; +use serde_json::{Map, Value}; + +use crate::{ + Channel, Client, ConnectTo, ConnectionTo, Dispatch, DynamicHandlerGuard, HandleDispatchFrom, + Handled, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, UntypedMessage, + role::{self}, + schema::v1, + util::MatchDispatchFrom, +}; + +#[doc(hidden)] +pub trait Wire: Send + Sync + 'static { + type Connect: JsonRpcRequest; + type Connected: JsonRpcResponse; + type Request: JsonRpcRequest; + type Notification: JsonRpcNotification; + type Response: JsonRpcResponse; + type Disconnect: JsonRpcRequest; + + fn connect(id: String) -> Self::Connect; + fn connected_id(response: Self::Connected) -> String; + fn request(id: String, method: String, params: Option>) -> Self::Request; + fn notification( + id: String, + method: String, + params: Option>, + ) -> Self::Notification; + fn incoming_request(request: Self::Request) -> (String, String, Option>); + fn incoming_notification( + notification: Self::Notification, + ) -> (String, String, Option>); + fn disconnect(id: String) -> Self::Disconnect; +} + +/// Stable ACP wire version. +#[derive(Debug)] +pub struct V1; +impl Wire for V1 { + type Connect = v1::ConnectMcpRequest; + type Connected = v1::ConnectMcpResponse; + type Request = v1::MessageMcpRequest; + type Notification = v1::MessageMcpNotification; + type Response = v1::MessageMcpResponse; + type Disconnect = v1::DisconnectMcpRequest; + + fn connect(id: String) -> Self::Connect { + v1::ConnectMcpRequest::new(v1::McpServerAcpId::new(id)) + } + fn connected_id(response: Self::Connected) -> String { + response.connection_id.0.to_string() + } + fn request(id: String, method: String, params: Option>) -> Self::Request { + v1::MessageMcpRequest::new(v1::McpConnectionId::new(id), method).params(params) + } + fn notification( + id: String, + method: String, + params: Option>, + ) -> Self::Notification { + v1::MessageMcpNotification::new(v1::McpConnectionId::new(id), method).params(params) + } + fn incoming_request(request: Self::Request) -> (String, String, Option>) { + ( + request.connection_id.0.to_string(), + request.method, + request.params, + ) + } + fn incoming_notification( + notification: Self::Notification, + ) -> (String, String, Option>) { + ( + notification.connection_id.0.to_string(), + notification.method, + notification.params, + ) + } + fn disconnect(id: String) -> Self::Disconnect { + v1::DisconnectMcpRequest::new(v1::McpConnectionId::new(id)) + } +} + +#[cfg(feature = "unstable_protocol_v2")] +/// Draft ACP v2 wire version. +#[derive(Debug)] +pub struct V2; +#[cfg(feature = "unstable_protocol_v2")] +impl Wire for V2 { + type Connect = crate::schema::v2::ConnectMcpRequest; + type Connected = crate::schema::v2::ConnectMcpResponse; + type Request = crate::schema::v2::MessageMcpRequest; + type Notification = crate::schema::v2::MessageMcpNotification; + type Response = crate::schema::v2::MessageMcpResponse; + type Disconnect = crate::schema::v2::DisconnectMcpRequest; + + fn connect(id: String) -> Self::Connect { + crate::schema::v2::ConnectMcpRequest::new(crate::schema::v2::McpServerAcpId::new(id)) + } + fn connected_id(response: Self::Connected) -> String { + response.connection_id.0.to_string() + } + fn request(id: String, method: String, params: Option>) -> Self::Request { + crate::schema::v2::MessageMcpRequest::new( + crate::schema::v2::McpConnectionId::new(id), + method, + ) + .params(params) + } + fn notification( + id: String, + method: String, + params: Option>, + ) -> Self::Notification { + crate::schema::v2::MessageMcpNotification::new( + crate::schema::v2::McpConnectionId::new(id), + method, + ) + .params(params) + } + fn incoming_request(request: Self::Request) -> (String, String, Option>) { + ( + request.connection_id.0.to_string(), + request.method, + request.params, + ) + } + fn incoming_notification( + notification: Self::Notification, + ) -> (String, String, Option>) { + ( + notification.connection_id.0.to_string(), + notification.method, + notification.params, + ) + } + fn disconnect(id: String) -> Self::Disconnect { + crate::schema::v2::DisconnectMcpRequest::new(crate::schema::v2::McpConnectionId::new(id)) + } +} + +fn params(value: Value) -> Result>, crate::Error> { + match value { + Value::Object(map) => Ok(Some(map)), + Value::Null => Ok(None), + _ => { + Err(crate::Error::invalid_params() + .data("native MCP parameters must be an object or null")) + } + } +} + +struct Incoming { + id: Arc>>, + tx: mpsc::UnboundedSender, + protocol: PhantomData

, +} + +impl HandleDispatchFrom for Incoming

{ + fn describe_chain(&self) -> impl std::fmt::Debug { + "McpOverAcp" + } + + async fn handle_dispatch_from( + &mut self, + message: Dispatch, + cx: ConnectionTo, + ) -> Result, crate::Error> { + MatchDispatchFrom::new(message, &cx) + .if_request_from(Client, async |request: P::Request, responder| { + let (id, method, params) = P::incoming_request(request); + if self.id.lock().unwrap().as_ref() != Some(&id) { + return Ok(Handled::No { + message: (P::request(id, method, params), responder), + retry: false, + }); + } + let responder = responder.wrap_params(|method, result| { + result.and_then(|value: Value| P::Response::from_value(method, value)) + }); + self.tx + .unbounded_send(Dispatch::Request( + UntypedMessage { + method, + params: params.map_or(Value::Null, Value::Object), + }, + responder, + )) + .map_err(crate::Error::into_internal_error)?; + Ok(Handled::Yes) + }) + .await + .if_notification_from(Client, async |notification: P::Notification| { + let (id, method, params) = P::incoming_notification(notification); + if self.id.lock().unwrap().as_ref() != Some(&id) { + return Ok(Handled::No { + message: P::notification(id, method, params), + retry: false, + }); + } + self.tx + .unbounded_send(Dispatch::Notification(UntypedMessage { + method, + params: params.map_or(Value::Null, Value::Object), + })) + .map_err(crate::Error::into_internal_error)?; + Ok(Handled::Yes) + }) + .await + .done() + } +} + +struct Closing { + cx: ConnectionTo, + id: String, + guard: Mutex>>, + closed: Arc, + disconnected: AtomicBool, + incoming: mpsc::UnboundedSender, + pending: Arc>>>, + protocol: PhantomData

, +} +impl Closing

{ + fn unregister(&self) -> bool { + let was_closed = self.closed.swap(true, Ordering::AcqRel); + self.guard.lock().unwrap().take(); + self.incoming.close_channel(); + let outstanding = std::mem::take(&mut *self.pending.lock().unwrap()); + for (_, responder) in outstanding { + drop(responder.respond_with_error(connection_closed())); + } + was_closed + } +} +fn connection_closed() -> crate::Error { + crate::Error::internal_error().data("MCP-over-ACP connection closed") +} +impl Drop for Closing

{ + fn drop(&mut self) { + self.unregister(); + if self.disconnected.load(Ordering::Acquire) { + return; + } + let cx = self.cx.clone(); + let id = self.id.clone(); + // A failed best-effort disconnect must not bring down the ACP connection. + drop(self.cx.spawn(async move { + drop( + cx.send_request_to(Client, P::disconnect(id)) + .block_task() + .await, + ); + Ok(()) + })); + } +} + +/// Cloneable shutdown handle. Call `close` for an acknowledged disconnect; +/// dropping the final handle only attempts a best-effort disconnect. +pub struct McpOverAcpClose(Arc>); +impl std::fmt::Debug for McpOverAcpClose

{ + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("McpOverAcpClose") + .field("id", &self.0.id) + .finish() + } +} +impl Clone for McpOverAcpClose

{ + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} +impl McpOverAcpClose

{ + /// Unregister callbacks immediately and await the provider's disconnect reply. + pub async fn close(self) -> Result<(), crate::Error> { + let already_closed = self.0.unregister(); + if already_closed { + return Ok(()); + } + let result = self + .0 + .cx + .send_request_to(Client, P::disconnect(self.0.id.clone())) + .block_task() + .await + .map(|_| ()); + if result.is_ok() { + self.0.disconnected.store(true, Ordering::Release); + } + result + } +} + +// Drain queued reverse requests even if the MCP component future is cancelled +// before its incoming task processes them. +struct IncomingQueue(mpsc::UnboundedReceiver); +impl Drop for IncomingQueue { + fn drop(&mut self) { + while let Ok(dispatch) = self.0.try_recv() { + if let Dispatch::Request(_, responder) = dispatch { + drop(responder.respond_with_error(connection_closed())); + } + } + } +} + +/// An MCP server transport backed by an already-open ACP `mcp/connect`. +/// +/// Pass this to `role::mcp::Client.builder().connect_with(...)`. Retain the +/// returned close handle to shut down while the MCP client is running. +pub struct McpOverAcp { + cx: ConnectionTo, + id: String, + rx: IncomingQueue, + close: McpOverAcpClose

, +} +impl std::fmt::Debug for McpOverAcp

{ + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("McpOverAcp") + .field("id", &self.id) + .finish_non_exhaustive() + } +} + +impl McpOverAcp { + /// Connect a v1 `McpServer::Acp` declaration from an ACP agent. + pub async fn connect_v1( + cx: &ConnectionTo, + id: v1::McpServerAcpId, + ) -> Result<(Self, McpOverAcpClose), crate::Error> { + connect::(cx, id.0.to_string()).await + } +} + +#[cfg(feature = "unstable_protocol_v2")] +impl McpOverAcp { + /// Connect a draft v2 `McpServer::Acp` declaration from an ACP agent. + pub async fn connect_v2( + cx: &crate::V2ConnectionTo, + id: crate::schema::v2::McpServerAcpId, + ) -> Result<(Self, McpOverAcpClose), crate::Error> { + connect::(cx.raw_connection(), id.0.to_string()).await + } +} + +async fn connect( + cx: &ConnectionTo, + server_id: String, +) -> Result<(McpOverAcp

, McpOverAcpClose

), crate::Error> { + let (tx, rx) = mpsc::unbounded(); + let id = Arc::new(Mutex::new(None)); + let guard = cx.add_dynamic_handler(Incoming::

{ + id: id.clone(), + tx: tx.clone(), + protocol: PhantomData, + })?; + // The registration is queued asynchronously. Apply it before publishing + // connect so peer callbacks cannot overtake the handler. + cx.dynamic_handler_barrier().await?; + let (reply_tx, reply_rx) = oneshot::channel(); + let reply_cx = cx.clone(); + cx.send_request_to(Client, P::connect(server_id)) + .on_receiving_result(move |reply| { + let reply = reply.map(P::connected_id); + if let Ok(connection_id) = &reply { + // The ACP response callback runs before the next incoming + // message, even when the peer sends an immediate MCP callback. + *id.lock().unwrap() = Some(connection_id.clone()); + } + if let Err(reply) = reply_tx.send(reply) { + // The caller cancelled while connect was outstanding. + if let Ok(connection_id) = reply { + let cx = reply_cx.clone(); + drop(reply_cx.spawn(async move { + drop( + cx.send_request_to(Client, P::disconnect(connection_id)) + .block_task() + .await, + ); + Ok(()) + })); + } + } + futures::future::ready(Ok(())) + })?; + let connection_id = reply_rx + .await + .map_err(crate::Error::into_internal_error)??; + let close = McpOverAcpClose(Arc::new(Closing { + cx: cx.clone(), + id: connection_id.clone(), + guard: Mutex::new(Some(guard)), + closed: Arc::new(AtomicBool::new(false)), + disconnected: AtomicBool::new(false), + incoming: tx, + pending: Arc::new(Mutex::new(HashMap::new())), + protocol: PhantomData, + })); + Ok(( + McpOverAcp { + cx: cx.clone(), + id: connection_id, + rx: IncomingQueue(rx), + close: close.clone(), + }, + close, + )) +} + +impl ConnectTo for McpOverAcp

{ + async fn connect_to( + self, + client: impl ConnectTo, + ) -> Result<(), crate::Error> { + let Self { + cx, + id, + mut rx, + close, + } = self; + let pending = close.0.pending.clone(); + let closed = close.0.closed.clone(); + let next_request = AtomicU64::new(0); + let (client_channel, server_channel) = Channel::duplex(); + let server = role::mcp::Server + .builder() + .on_receive_dispatch( + async move |message: Dispatch, _mcp_cx| match message { + Dispatch::Request(request, responder) => { + let (method, value) = request.into_parts(); + let params = match params(value) { + Ok(params) => params, + Err(error) => return responder.respond_with_error(error), + }; + let request = P::request(id.clone(), method, params); + let responder = responder.wrap_params(|method, result| { + result.and_then(|response: P::Response| response.into_json(method)) + }); + cx.send_proxied_message_to( + Client, + Dispatch::::Request(request, responder), + ) + } + Dispatch::Notification(notification) => { + let (method, value) = notification.into_parts(); + let params = params(value)?; + cx.send_notification_to(Client, P::notification(id.clone(), method, params)) + } + Dispatch::Response(result, router) => router.route_with_result(result), + }, + crate::on_receive_dispatch!(), + ) + .with_spawned(move |mcp_cx| async move { + while let Some(message) = rx.0.next().await { + match message { + Dispatch::Request(request, responder) => { + let mut outstanding = pending.lock().unwrap(); + if closed.load(Ordering::Acquire) { + drop(outstanding); + responder.respond_with_error(connection_closed())?; + continue; + } + let sequence = next_request.fetch_add(1, Ordering::Relaxed); + outstanding.insert(sequence, responder); + drop(outstanding); + let waiting = mcp_cx.send_request_to(role::mcp::Client, request); + let pending = pending.clone(); + mcp_cx.spawn(async move { + let result = waiting.block_task().await; + let responder = pending.lock().unwrap().remove(&sequence); + if let Some(responder) = responder { + // The peer may have closed while the MCP request + // was in flight; no error escapes to the ACP loop. + drop(responder.respond_with_result(result)); + } + Ok(()) + })?; + } + message => mcp_cx.send_proxied_message_to(role::mcp::Client, message)?, + } + } + Ok(()) + }); + futures::try_join!( + server.connect_to(server_channel), + client.connect_to(client_channel) + )?; + Ok(()) + } +} 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..f768a29b 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -1,7 +1,11 @@ -use std::{marker::PhantomData, sync::Arc}; +use std::{ + marker::PhantomData, + sync::{Arc, Mutex, Weak}, +}; -use futures::channel::mpsc; -use futures::{SinkExt, StreamExt}; +use futures::channel::{mpsc, oneshot}; +use futures::future::{Either, Shared, select}; +use futures::{FutureExt, SinkExt, StreamExt}; use rustc_hash::FxHashMap; use serde_json::{Map, Value}; @@ -205,11 +209,34 @@ pub(super) struct McpActiveSession mcp_connect: Arc>, /// Active connections to MCP server tasks. - connections: FxHashMap>, + connections: Arc>>, protocol: PhantomData Protocol>, } +struct NativeConnection { + sender: mpsc::Sender, + shutdown: oneshot::Sender<()>, + finished: Shared>, +} + +struct NativeConnectionGuard { + id: McpConnectionId, + connections: Weak>>, + finished: Option>, +} + +impl Drop for NativeConnectionGuard { + fn drop(&mut self) { + if let Some(connections) = self.connections.upgrade() { + connections.lock().unwrap().remove(&self.id); + } + if let Some(finished) = self.finished.take() { + let _ = finished.send(()); + } + } +} + impl McpActiveSession where Counterpart: HasPeer, @@ -222,7 +249,7 @@ where Self { server_id, mcp_connect, - connections: FxHashMap::default(), + connections: Arc::new(Mutex::new(FxHashMap::default())), protocol: PhantomData, } } @@ -251,14 +278,17 @@ where 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 (shutdown, shutdown_rx) = oneshot::channel(); + let (finished_tx, finished) = oneshot::channel(); + let finished = finished.shared(); let (client_channel, server_channel) = Channel::duplex(); let client_component = { let connection_id = connection_id.clone(); let acp_connection = acp_connection.clone(); + let relay_connection = acp_connection.clone(); + let relay_finished = finished.clone(); role::mcp::Client .builder() @@ -313,7 +343,28 @@ where .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)?; + match message { + Dispatch::Request(request, responder) => { + let response = mcp_connection + .send_request_to(role::mcp::Server, request) + .block_task(); + let finished = relay_finished.clone(); + // Keep the response waiter on ACP, not on the + // child actor being stopped. A pending request + // receives an error when that child closes. + relay_connection.spawn(async move { + let result = match select(Box::pin(response), finished).await { + Either::Left((result, _)) => result, + Either::Right(_) => Err(crate::Error::internal_error() + .data("native MCP connection closed")), + }; + responder.respond_with_result(result) + })?; + } + other => { + mcp_connection.send_proxied_message_to(role::mcp::Server, other)? + } + } } Ok(()) }) @@ -327,19 +378,43 @@ where connection: acp_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 }) - }); + let guard = NativeConnectionGuard { + id: connection_id.clone(), + connections: Arc::downgrade(&self.connections), + finished: Some(finished_tx), + }; + // Both halves belong to one task. A child error must not propagate + // through the task actor and terminate the parent ACP connection. + let spawn_result = acp_connection.spawn(async move { + let client = Box::pin(client_component.connect_to(client_channel)); + let server = Box::pin(spawned_server.connect_to(server_channel)); + let running = select(client, server); + let result = match select(Box::pin(running), shutdown_rx).await { + Either::Left((Either::Left((result, _)), _)) + | Either::Left((Either::Right((result, _)), _)) => Some(result), + Either::Right(_) => None, + }; + if let Some(Err(error)) = result { + tracing::warn!(?error, "native MCP connection closed with error"); + } + drop(guard); + Ok(()) + }); - match spawn_results { + match spawn_result { Ok(()) => { + self.connections.lock().unwrap().insert( + connection_id.clone(), + NativeConnection { + sender: mcp_server_tx, + shutdown, + finished, + }, + ); 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) } @@ -359,7 +434,13 @@ where crate::Error, > { let connection_id = Protocol::message_request_connection_id(&request); - let Some(mcp_server_tx) = self.connections.get_mut(&connection_id) else { + let sender = self + .connections + .lock() + .unwrap() + .get(&connection_id) + .map(|entry| entry.sender.clone()); + let Some(mut mcp_server_tx) = sender else { return Ok(Handled::No { message: (request, responder), retry: false, @@ -375,10 +456,14 @@ where result .and_then(|response: Value| Protocol::MessageResponse::from_value(method, response)) }); - mcp_server_tx + if let Err(error) = mcp_server_tx .send(Dispatch::Request(untyped, responder)) .await - .map_err(crate::Error::into_internal_error)?; + { + // Dropping the undelivered responder reports failure to the + // caller; a closed child must not terminate its parent ACP task. + tracing::debug!(?error, "native MCP relay closed during request"); + } Ok(Handled::Yes) } @@ -389,7 +474,13 @@ where 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 { + let sender = self + .connections + .lock() + .unwrap() + .get(&connection_id) + .map(|entry| entry.sender.clone()); + let Some(mut mcp_server_tx) = sender else { return Ok(Handled::No { message: notification, retry: false, @@ -401,16 +492,15 @@ where 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)?; + if let Err(error) = mcp_server_tx.send(Dispatch::Notification(untyped)).await { + tracing::debug!(?error, "native MCP relay closed during notification"); + } Ok(Handled::Yes) } /// Disconnect an active native MCP-over-ACP connection. - fn handle_mcp_disconnect_request( + async fn handle_mcp_disconnect_request( &mut self, request: Protocol::DisconnectRequest, responder: Responder, @@ -422,13 +512,20 @@ where crate::Error, > { let connection_id = Protocol::disconnect_connection_id(&request); - if self.connections.remove(&connection_id).is_none() { + let entry = self.connections.lock().unwrap().remove(&connection_id); + let Some(NativeConnection { + shutdown, finished, .. + }) = entry + else { return Ok(Handled::No { message: (request, responder), retry: false, }); - } + }; + // A successful disconnect means both child halves have been dropped. + let _ = shutdown.send(()); + let _ = finished.await; responder.respond(Protocol::disconnect_response())?; Ok(Handled::Yes) } @@ -474,7 +571,7 @@ where .if_request_from( Agent, async |request: Protocol::DisconnectRequest, responder| { - self.handle_mcp_disconnect_request(request, responder) + self.handle_mcp_disconnect_request(request, responder).await }, ) .await diff --git a/src/agent-client-protocol/tests/mcp_connection_lifecycle.rs b/src/agent-client-protocol/tests/mcp_connection_lifecycle.rs new file mode 100644 index 00000000..a3303e45 --- /dev/null +++ b/src/agent-client-protocol/tests/mcp_connection_lifecycle.rs @@ -0,0 +1,545 @@ +#![cfg(all(feature = "unstable_mcp_over_acp", feature = "unstable_protocol_v2"))] + +use std::{future::pending, time::Duration}; + +use agent_client_protocol::{ + Agent, Client, ConnectTo, ConnectionTo, DynConnectTo, Error, JsonRpcRequest, JsonRpcResponse, + NullRun, Responder, UntypedMessage, V2ConnectionTo, + mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, + role, + schema::{ProtocolVersion, v1, v2}, +}; +use futures::{ + StreamExt, + channel::{mpsc, oneshot}, +}; +use serde::{Deserialize, Serialize}; +use serde_json::{Value, json}; + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)] +#[request(method = "_test/echo", response = EchoResponse)] +struct EchoRequest { + value: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)] +struct EchoResponse { + value: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)] +#[request(method = "_test/pending", response = EchoResponse)] +struct PendingRequest {} + +struct Probe { + id: String, + dropped: mpsc::UnboundedSender, +} + +impl Drop for Probe { + fn drop(&mut self) { + drop(self.dropped.unbounded_send(self.id.clone())); + } +} + +struct Server { + dropped: mpsc::UnboundedSender, + failures: mpsc::UnboundedSender<(String, oneshot::Sender<()>)>, + pending_started: mpsc::UnboundedSender<()>, + reverse_seen: mpsc::UnboundedSender, +} + +impl McpServerConnect for Server { + fn name(&self) -> String { + "lifecycle-test".into() + } + + fn connect(&self, context: McpConnectionTo) -> DynConnectTo { + let id = context.connection_id().unwrap().to_string(); + let (fail, failure) = oneshot::channel(); + self.failures.unbounded_send((id.clone(), fail)).unwrap(); + DynConnectTo::new(ServerConnection { + probe: Probe { + id, + dropped: self.dropped.clone(), + }, + failure, + pending_started: self.pending_started.clone(), + reverse_seen: self.reverse_seen.clone(), + }) + } +} + +struct ServerConnection { + probe: Probe, + failure: oneshot::Receiver<()>, + pending_started: mpsc::UnboundedSender<()>, + reverse_seen: mpsc::UnboundedSender, +} + +impl ConnectTo for ServerConnection { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + let pending_started = self.pending_started; + role::mcp::Server + .builder() + .on_receive_request( + async |request: EchoRequest, responder: Responder, _connection| { + responder.respond(EchoResponse { + value: request.value, + }) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |_request: PendingRequest, + _responder: Responder, + _connection| { + pending_started + .unbounded_send(()) + .map_err(Error::into_internal_error)?; + pending::<()>().await; + Ok(()) + }, + agent_client_protocol::on_receive_request!(), + ) + .with_spawned(move |connection| async move { + let _probe = self.probe; + connection.send_notification(UntypedMessage { + method: "_test/reverse-notice".into(), + params: Value::Null, + })?; + let response = connection + .send_request(EchoRequest { + value: "reverse".into(), + }) + .block_task() + .await?; + assert_eq!(response.value, "reverse"); + assert!( + connection + .send_request(EchoRequest { + value: "error".into() + }) + .block_task() + .await + .is_err() + ); + self.reverse_seen + .unbounded_send(_probe.id.clone()) + .map_err(Error::into_internal_error)?; + self.failure.await.map_err(Error::into_internal_error)?; + Err(Error::internal_error().data("child failure")) + }) + .connect_to(client) + .await + } +} + +fn implementation() -> v2::Implementation { + v2::Implementation::new("lifecycle-test", env!("CARGO_PKG_VERSION")) +} + +async fn echo( + connection: &V2ConnectionTo, + id: v2::McpConnectionId, + value: &str, +) -> Result<(), Error> { + let response = connection + .send_request( + v2::MessageMcpRequest::new(id, "_test/echo").params( + serde_json::from_value::>(json!({"value": value})) + .unwrap(), + ), + ) + .block_task() + .await?; + let response: Value = + serde_json::from_str(response.0.get()).map_err(Error::into_internal_error)?; + assert_eq!(response, json!({"value": value})); + Ok(()) +} + +async fn scenario( + connection: V2ConnectionTo, + server: v2::McpServerAcpId, + mut dropped: mpsc::UnboundedReceiver, + mut failures: mpsc::UnboundedReceiver<(String, oneshot::Sender<()>)>, + mut pending_started: mpsc::UnboundedReceiver<()>, + mut reverse_seen: mpsc::UnboundedReceiver, + mut reverse_notices: mpsc::UnboundedReceiver, +) -> Result<(), Error> { + let a = connection + .send_request(v2::ConnectMcpRequest::new(server.clone())) + .block_task() + .await? + .connection_id; + let b = connection + .send_request(v2::ConnectMcpRequest::new(server)) + .block_task() + .await? + .connection_id; + assert_ne!(a, b); + let (id_a, _fail_a) = failures.next().await.unwrap(); + let (id_b, fail_b) = failures.next().await.unwrap(); + assert_eq!((id_a, id_b), (a.to_string(), b.to_string())); + assert_eq!(reverse_seen.next().await, Some(a.to_string())); + assert_eq!(reverse_seen.next().await, Some(b.to_string())); + assert_eq!(reverse_notices.next().await, Some(a.to_string())); + assert_eq!(reverse_notices.next().await, Some(b.to_string())); + echo(&connection, a.clone(), "A").await?; + echo(&connection, b.clone(), "B").await?; + + // Leave a request outstanding when A disconnects. Neither that request + // nor the other half of A may keep its server alive. + let (pending_result_tx, pending_result_rx) = oneshot::channel(); + let pending_connection = connection.clone(); + let pending_id = a.clone(); + connection.spawn(async move { + let result = pending_connection + .send_request( + v2::MessageMcpRequest::new(pending_id, "_test/pending") + .params(serde_json::Map::new()), + ) + .block_task() + .await; + drop(pending_result_tx.send(result)); + Ok(()) + })?; + pending_started + .next() + .await + .expect("pending request reached server"); + connection + .send_request(v2::DisconnectMcpRequest::new(a.clone())) + .block_task() + .await?; + assert_eq!(dropped.next().await, Some(a.to_string())); + assert!( + pending_result_rx.await.unwrap().is_err(), + "pending request must fail on disconnect" + ); + assert!( + connection + .send_request(v2::DisconnectMcpRequest::new(a.clone())) + .block_task() + .await + .is_err() + ); + assert!( + connection + .send_request(v2::DisconnectMcpRequest::new(v2::McpConnectionId::new( + "missing" + ))) + .block_task() + .await + .is_err() + ); + assert!( + connection + .send_request(v2::MessageMcpRequest::new(a, "_test/echo")) + .block_task() + .await + .is_err() + ); + echo(&connection, b.clone(), "still alive").await?; + fail_b.send(()).unwrap(); + assert_eq!(dropped.next().await, Some(b.to_string())); + assert!( + connection + .send_request(v2::DisconnectMcpRequest::new(b)) + .block_task() + .await + .is_err() + ); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn native_connections_close_independently_without_closing_acp() -> Result<(), Error> { + let test = async { + let (server_tx, mut server_rx) = mpsc::unbounded(); + let (dropped_tx, dropped_rx) = mpsc::unbounded(); + let (failure_tx, failure_rx) = mpsc::unbounded(); + let (pending_started_tx, pending_started_rx) = mpsc::unbounded(); + let (reverse_seen_tx, reverse_seen_rx) = mpsc::unbounded(); + let (reverse_notice_tx, reverse_notice_rx) = mpsc::unbounded(); + let (result_tx, mut result_rx) = mpsc::unbounded(); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _connection: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new(request.protocol_version, implementation()) + .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, + _connection: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("missing native server") + }; + server_tx + .unbounded_send(server.server_id.clone()) + .map_err(Error::into_internal_error)?; + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new("lifecycle"))) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::MessageMcpRequest, + responder: Responder, + _connection: V2ConnectionTo| { + assert_eq!(request.method, "_test/echo"); + let value = request + .params + .as_ref() + .and_then(|params| params.get("value")); + if value == Some(&json!("error")) { + return responder.respond_with_error(Error::invalid_params()); + } + assert_eq!(value, Some(&json!("reverse"))); + let raw = serde_json::value::to_raw_value(&json!({"value": "reverse"})) + .map_err(Error::into_internal_error)?; + responder.respond(v2::MessageMcpResponse::new(raw.into())) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_notification( + async move |notification: v2::MessageMcpNotification, + _connection: V2ConnectionTo| { + assert_eq!(notification.method, "_test/reverse-notice"); + assert!( + notification.params.is_none(), + "null params must stay omitted" + ); + reverse_notice_tx + .unbounded_send(notification.connection_id.to_string()) + .map_err(Error::into_internal_error) + }, + agent_client_protocol::on_receive_notification!(), + ) + .with_spawned(move |connection: V2ConnectionTo| async move { + let server = server_rx.next().await.unwrap(); + let result = scenario( + connection, + server, + dropped_rx, + failure_rx, + pending_started_rx, + reverse_seen_rx, + reverse_notice_rx, + ) + .await; + result_tx + .unbounded_send(result) + .map_err(Error::into_internal_error) + }); + + Client + .v2() + .connect_with(agent, async move |connection| { + connection + .send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + implementation(), + )) + .block_task() + .await?; + let session = connection + .build_session_from(v2::NewSessionRequest::new( + std::env::current_dir().map_err(Error::into_internal_error)?, + )) + .with_mcp_server(McpServer::::new( + Server { + dropped: dropped_tx, + failures: failure_tx, + pending_started: pending_started_tx, + reverse_seen: reverse_seen_tx, + }, + NullRun, + ))? + .start_session() + .block_task() + .await?; + let result = result_rx + .next() + .await + .expect("agent scenario did not complete"); + drop(session); + result + }) + .await + }; + tokio::time::timeout(Duration::from_secs(10), test) + .await + .expect("native lifecycle timed out") +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn v1_native_connections_share_the_same_isolated_lifecycle() -> Result<(), Error> { + let test = async { + let (server_tx, mut server_rx) = mpsc::unbounded(); + let (dropped_tx, mut dropped_rx) = mpsc::unbounded(); + let (failure_tx, mut failure_rx) = mpsc::unbounded::<(String, oneshot::Sender<()>)>(); + let (reverse_seen_tx, mut reverse_seen_rx) = mpsc::unbounded(); + let (reverse_notice_tx, mut reverse_notice_rx) = mpsc::unbounded(); + let (result_tx, mut result_rx) = mpsc::unbounded(); + let (pending_started_tx, _pending_started_rx) = mpsc::unbounded(); + let agent = Agent + .builder() + .on_receive_request( + async move |request: v1::NewSessionRequest, + responder: Responder, + _connection: ConnectionTo| { + let [v1::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("missing v1 native server") + }; + server_tx + .unbounded_send(server.server_id.clone()) + .map_err(Error::into_internal_error)?; + responder.respond(v1::NewSessionResponse::new("v1-lifecycle")) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v1::MessageMcpRequest, + responder: Responder, + _connection: ConnectionTo| { + assert_eq!(request.method, "_test/echo"); + if request + .params + .as_ref() + .and_then(|params| params.get("value")) + == Some(&json!("error")) + { + return responder.respond_with_error(Error::invalid_params()); + } + let raw = serde_json::value::to_raw_value(&json!({"value": "reverse"})) + .map_err(Error::into_internal_error)?; + responder.respond(v1::MessageMcpResponse::new(raw.into())) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_notification( + async move |notification: v1::MessageMcpNotification, + _connection: ConnectionTo| { + assert_eq!(notification.method, "_test/reverse-notice"); + assert!(notification.params.is_none()); + reverse_notice_tx + .unbounded_send(notification.connection_id.to_string()) + .map_err(Error::into_internal_error) + }, + agent_client_protocol::on_receive_notification!(), + ) + .with_spawned(move |connection: ConnectionTo| async move { + let server = server_rx.next().await.unwrap(); + let result = async { + let a = connection + .send_request(v1::ConnectMcpRequest::new(server.clone())) + .block_task() + .await? + .connection_id; + let b = connection + .send_request(v1::ConnectMcpRequest::new(server)) + .block_task() + .await? + .connection_id; + assert_ne!(a, b); + let (id_a, _failure_a) = failure_rx.next().await.unwrap(); + let (id_b, failure_b) = failure_rx.next().await.unwrap(); + assert_eq!((id_a, id_b), (a.to_string(), b.to_string())); + assert_eq!(reverse_seen_rx.next().await, Some(a.to_string())); + assert_eq!(reverse_seen_rx.next().await, Some(b.to_string())); + assert_eq!(reverse_notice_rx.next().await, Some(a.to_string())); + assert_eq!(reverse_notice_rx.next().await, Some(b.to_string())); + let response = connection + .send_request( + v1::MessageMcpRequest::new(a.clone(), "_test/echo").params( + serde_json::Map::from_iter([("value".into(), json!("v1"))]), + ), + ) + .block_task() + .await?; + assert_eq!( + serde_json::from_str::(response.0.get()).unwrap(), + json!({"value": "v1"}) + ); + connection + .send_request(v1::DisconnectMcpRequest::new(a.clone())) + .block_task() + .await?; + assert_eq!(dropped_rx.next().await, Some(a.to_string())); + assert!( + connection + .send_request(v1::DisconnectMcpRequest::new(a)) + .block_task() + .await + .is_err() + ); + let response = connection + .send_request(v1::MessageMcpRequest::new(b.clone(), "_test/echo").params( + serde_json::Map::from_iter([("value".into(), json!("still open"))]), + )) + .block_task() + .await?; + assert_eq!( + serde_json::from_str::(response.0.get()).unwrap(), + json!({"value": "still open"}) + ); + failure_b.send(()).unwrap(); + assert_eq!(dropped_rx.next().await, Some(b.to_string())); + assert!( + connection + .send_request(v1::DisconnectMcpRequest::new(b)) + .block_task() + .await + .is_err() + ); + Ok::<(), Error>(()) + } + .await; + result_tx + .unbounded_send(result) + .map_err(Error::into_internal_error) + }); + + Client + .builder() + .connect_with(agent, async move |connection| { + let session = connection + .build_session_cwd()? + .with_mcp_server(McpServer::::new( + Server { + dropped: dropped_tx, + failures: failure_tx, + pending_started: pending_started_tx, + reverse_seen: reverse_seen_tx, + }, + NullRun, + ))? + .block_task() + .start_session() + .await?; + let result = result_rx + .next() + .await + .expect("v1 scenario did not complete"); + drop(session); + result + }) + .await + }; + tokio::time::timeout(Duration::from_secs(10), test) + .await + .expect("v1 native lifecycle timed out") +} diff --git a/src/agent-client-protocol/tests/native_mcp_consumer.rs b/src/agent-client-protocol/tests/native_mcp_consumer.rs new file mode 100644 index 00000000..5a296704 --- /dev/null +++ b/src/agent-client-protocol/tests/native_mcp_consumer.rs @@ -0,0 +1,319 @@ +#![cfg(feature = "unstable_mcp_over_acp")] + +use std::{ + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, Channel, Client, Error, Responder, + mcp_client::McpOverAcp, + role, + schema::v1::{ + ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, DisconnectMcpResponse, + McpConnectionId, McpServerAcpId, MessageMcpNotification, MessageMcpRequest, + MessageMcpResponse, + }, +}; +use futures::channel::oneshot; +use serde_json::{Value, json}; + +#[tokio::test] +async fn native_mcp_initialize_tools_callback_and_close() { + let closed = Arc::new(AtomicBool::new(false)); + let closed_provider = closed.clone(); + let (reverse_result_tx, reverse_result_rx) = oneshot::channel(); + let reverse_result_tx = Arc::new(Mutex::new(Some(reverse_result_tx))); + let (reverse_received_tx, reverse_received_rx) = oneshot::channel(); + let reverse_received_tx = Mutex::new(Some(reverse_received_tx)); + let held_reverse_responder = Arc::new(Mutex::new(None::>)); + let held_for_handler = held_reverse_responder.clone(); + let (channel_agent, channel_client) = Channel::duplex(); + let provider = Client.builder() + .on_receive_request( + async |request: ConnectMcpRequest, responder: Responder, cx| { + assert_eq!(request.server_id.0.as_ref(), "demo-server"); + responder.respond(ConnectMcpResponse::new(McpConnectionId::new("demo-connection")))?; + cx.send_notification(MessageMcpNotification::new( + McpConnectionId::new("demo-connection"), "notifications/tools/list_changed", + ))?; + Ok(()) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: MessageMcpRequest, responder: Responder, cx| { + assert_eq!(request.connection_id.0.as_ref(), "demo-connection"); + let result = match request.method.as_str() { + "initialize" => { + assert!(request.params.is_some()); + json!({"protocolVersion":"2024-11-05","capabilities":{"tools":{}},"serverInfo":{"name":"demo","version":"1"}}) + } + "tools/list" => { + assert!(request.params.is_none()); + let reverse_result_tx = reverse_result_tx.clone(); + cx.send_request(MessageMcpRequest::new( + McpConnectionId::new("demo-connection"), "sampling/createMessage", + )).on_receiving_result(move |result| { + if let Some(tx) = reverse_result_tx.lock().unwrap().take() { + let _ = tx.send(result.is_err()); + } + futures::future::ready(Ok(())) + })?; + json!({"tools":[{"name":"echo","description":"Echo","inputSchema":{"type":"object"}}]}) + } + other => panic!("unexpected MCP request {other}"), + }; + responder.respond(agent_client_protocol::JsonRpcResponse::from_value("mcp/message", result)?)?; + Ok(()) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: DisconnectMcpRequest, responder: Responder, _| { + assert_eq!(request.connection_id.0.as_ref(), "demo-connection"); + closed_provider.store(true, Ordering::Release); + responder.respond(DisconnectMcpResponse::new()) + }, + agent_client_protocol::on_receive_request!(), + ); + let agent = Agent.builder().connect_with(channel_agent, async |cx| { + let (transport, close) = + McpOverAcp::connect_v1(&cx, McpServerAcpId::new("demo-server")).await?; + let (notice_tx, notice_rx) = oneshot::channel(); + let notice_tx = Mutex::new(Some(notice_tx)); + let mcp_client = role::mcp::Client + .builder() + .on_receive_notification( + async move |notice: agent_client_protocol::UntypedMessage, _| { + assert_eq!(notice.method(), "notifications/tools/list_changed"); + if let Some(tx) = notice_tx.lock().unwrap().take() { + let _ = tx.send(()); + } + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ) + .on_receive_request( + async move |request: agent_client_protocol::UntypedMessage, + responder: Responder, + _| { + assert_eq!(request.method(), "sampling/createMessage"); + *held_for_handler.lock().unwrap() = Some(responder); + if let Some(tx) = reverse_received_tx.lock().unwrap().take() { + let _ = tx.send(()); + } + Ok(()) + }, + agent_client_protocol::on_receive_request!(), + ); + mcp_client + .connect_with(transport, async |mcp| { + let init = mcp + .send_request(agent_client_protocol::UntypedMessage { + method: "initialize".into(), + params: json!({ + "protocolVersion":"2024-11-05", + "capabilities":{}, + "clientInfo":{"name":"consumer","version":"1"} + }), + }) + .block_task() + .await?; + assert_eq!(init["serverInfo"]["name"], "demo"); + let listed = mcp + .send_request(agent_client_protocol::UntypedMessage { + method: "tools/list".into(), + params: Value::Null, + }) + .block_task() + .await?; + assert_eq!(listed["tools"][0]["name"], "echo"); + notice_rx.await.map_err(Error::into_internal_error)?; + reverse_received_rx + .await + .map_err(Error::into_internal_error)?; + close.close().await?; + assert!( + reverse_result_rx + .await + .map_err(Error::into_internal_error)?, + "pending reverse request should fail when MCP transport closes" + ); + Ok(()) + }) + .await?; + Ok(()) + }); + tokio::time::timeout(Duration::from_secs(10), async { + futures::try_join!(provider.connect_to(channel_client), agent) + }) + .await + .expect("native MCP connection timed out") + .expect("native MCP connection failed"); + assert!( + closed.load(Ordering::Acquire), + "provider did not receive disconnect" + ); + drop(held_reverse_responder.lock().unwrap().take()); +} + +#[cfg(feature = "unstable_protocol_v2")] +#[tokio::test] +async fn draft_v2_uses_the_same_native_mcp_transport() { + use agent_client_protocol::{ + mcp_client::V2, + schema::{ProtocolVersion, v2}, + }; + + let (agent_channel, client_channel) = Channel::duplex(); + let (initialized_tx, initialized_rx) = oneshot::channel(); + let provider = Client + .v2() + .on_receive_request( + async |request: v2::ConnectMcpRequest, + responder: Responder, + _| { + assert_eq!(request.server_id.0.as_ref(), "v2-server"); + responder.respond(v2::ConnectMcpResponse::new("v2-connection")) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::MessageMcpRequest, + responder: Responder, + _| { + assert_eq!(request.connection_id.0.as_ref(), "v2-connection"); + assert_eq!(request.method, "tools/list"); + assert!(request.params.is_none()); + responder.respond(agent_client_protocol::JsonRpcResponse::from_value( + "mcp/message", + json!({"tools":[{"name":"v2-echo","inputSchema":{"type":"object"}}]}), + )?) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::DisconnectMcpRequest, + responder: Responder, + _| { + assert_eq!(request.connection_id.0.as_ref(), "v2-connection"); + responder.respond(v2::DisconnectMcpResponse::new()) + }, + agent_client_protocol::on_receive_request!(), + ); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("consumer-test-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new() + .session(v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + )), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .connect_with(agent_channel, async |cx| { + initialized_rx.await.map_err(Error::into_internal_error)?; + let (transport, close) = + McpOverAcp::::connect_v2(&cx, v2::McpServerAcpId::new("v2-server")).await?; + role::mcp::Client + .builder() + .connect_with(transport, async |mcp| { + let result = mcp + .send_request(agent_client_protocol::UntypedMessage { + method: "tools/list".into(), + params: Value::Null, + }) + .block_task() + .await?; + assert_eq!(result["tools"][0]["name"], "v2-echo"); + close.close().await + }) + .await + }); + let provider = provider.connect_with(client_channel, async |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("consumer-test-client", "1"), + )) + .block_task() + .await?; + let _ = initialized_tx.send(()); + cx.incoming_closed().await; + Ok(()) + }); + tokio::time::timeout(Duration::from_secs(10), async { + futures::try_join!(provider, agent) + }) + .await + .expect("v2 MCP connection timed out") + .expect("v2 MCP connection failed"); +} + +#[tokio::test] +async fn cancelled_connect_disconnects_if_provider_opens_later() { + let (held_tx, held_rx) = oneshot::channel(); + let held_tx = Mutex::new(Some(held_tx)); + let (disconnected_tx, disconnected_rx) = oneshot::channel(); + let disconnected_tx = Mutex::new(Some(disconnected_tx)); + let (agent_channel, client_channel) = Channel::duplex(); + let provider = Client + .builder() + .on_receive_request( + async move |_: ConnectMcpRequest, responder: Responder, _| { + if let Some(tx) = held_tx.lock().unwrap().take() { + let _ = tx.send(responder); + } + Ok(()) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: DisconnectMcpRequest, + responder: Responder, + _| { + assert_eq!(request.connection_id.0.as_ref(), "late-connection"); + if let Some(tx) = disconnected_tx.lock().unwrap().take() { + let _ = tx.send(()); + } + responder.respond(DisconnectMcpResponse::new()) + }, + agent_client_protocol::on_receive_request!(), + ); + let agent = Agent.builder().connect_with(agent_channel, async |cx| { + let mut pending = Box::pin(McpOverAcp::connect_v1( + &cx, + McpServerAcpId::new("late-server"), + )); + let responder = tokio::select! { + result = &mut pending => panic!("connect completed before provider responded: {result:?}"), + result = held_rx => result.map_err(Error::into_internal_error)?, + }; + drop(pending); + responder.respond(ConnectMcpResponse::new(McpConnectionId::new( + "late-connection", + )))?; + disconnected_rx.await.map_err(Error::into_internal_error)?; + Ok(()) + }); + tokio::time::timeout(Duration::from_secs(10), async { + futures::try_join!(provider.connect_to(client_channel), agent) + }) + .await + .expect("cancelled connect cleanup timed out") + .expect("cancelled connect cleanup failed"); +}